Loading...

verticapy.machine_learning.memmodel.tree.NonBinaryTree

class verticapy.machine_learning.memmodel.tree.NonBinaryTree(tree: dict, classes: Annotated[list | ndarray, 'Array Like Structure'] | None = None)

InMemoryModel implementation of non-binary trees.

Parameters

tree: dict

A NonBinaryTree tree. NonBinaryTree can be generated with the vDataFrame.chaid() method.

classes: ArrayLike, optional

The classes for the non-binary tree model.

Attributes

Attributes are identical to the input parameters, followed by an underscore (‘_’).

Examples

Initalization

Import the required module.

from verticapy.machine_learning.memmodel.tree import NonBinaryTree

A NonBinaryTree tree model is defined by the non-binary decision tree and name of classes.

We will first generate a non-binary tree using vDataFrame.chaid() method. For this example, we will use the Titanic dataset.

import verticapy.datasets as vpd

data = vpd.load_titanic()
123
pclass
Integer
123
survived
Integer
Abc
Varchar(164)
Abc
sex
Varchar(20)
123
age
Numeric(8)
123
sibsp
Integer
123
parch
Integer
Abc
ticket
Varchar(36)
123
fare
Numeric(12)
Abc
cabin
Varchar(30)
Abc
embarked
Varchar(20)
Abc
boat
Varchar(100)
123
body
Integer
Abc
Varchar(100)
110male71.000PC 1760949.5042[null]C[null]22
210male45.00011378435.5TS[null][null]
310male[null]0011379831.0[null]S[null][null]
410male17.00011305947.1[null]S[null][null]
510male27.01013508136.7792C89C[null][null]
610male37.011PC 1775683.1583E52C[null][null]
710male31.010F.C. 1275052.0B71S[null][null]
810male50.010PC 17761106.425C86C[null]62
910female36.000PC 1753131.6792A29C[null][null]
1010male37.01011380353.1C123S[null][null]
1110male24.000PC 1759379.2B86C[null][null]
1210male45.0103697383.475C83S[null][null]
1310male40.0001120590.0B94S[null]110
1410male42.00011303842.5B11S[null][null]
1510male[null]001746351.8625E46S[null][null]
1610male42.01011378952.0[null]S[null]38
1710male[null]00PC 1760030.6958[null]C14[null]
1810male29.00011350130.0D6S[null]126
1910male46.0001305075.2417C6C[null]292
2010male54.0001746351.8625E46S[null]175
2110male47.00011379642.4[null]S[null][null]
2210male58.00235273113.275D48C[null]122
2310male45.50011304328.5C124S[null]166
2410male29.01011377666.6C2S[null][null]
2510male47.00011046552.0C110S[null]207
2610male38.000199720.0[null]S[null][null]
2710male22.000PC 17760135.6333[null]C[null]232
2810male31.000PC 1759050.4958A24S[null][null]
2910male50.0101350755.9E44S[null][null]
3010male56.0001776430.6958A7C[null][null]
3110male57.010PC 17569146.5208B78C[null][null]
3210female63.010PC 17483221.7792C55 C57S[null][null]
3310male61.0003696332.3208D50S[null]46
3410male21.0013528177.2875D26S[null]169
3510male51.001PC 1759761.3792[null]C[null][null]
3611female63.0101350277.9583D7S10[null]
3711female32.0001181376.2917D15C8[null]
3811female58.00011378326.55C103S8[null]
3911female44.000PC 1761027.7208B4C6[null]
4011female41.00016966134.5E40C3[null]
4111female53.000PC 1760627.4458[null]C6[null]
4211male36.001PC 17755512.3292B51 B53 B55C3[null]
4311female58.001PC 17755512.3292B51 B53 B55C3[null]
4411male11.012113760120.0B96 B98S4[null]
4511female76.0101987778.85C46S6[null]
4611female[null]0111350555.0E33S6[null]
4711female39.011PC 1775683.1583E49C14[null]
4811female27.012F.C. 1275052.0B71S3[null]
4911female[null]0017421110.8833[null]C4[null]
5011female35.000113503211.5C130C4[null]
5111female22.00111237859.4[null]C7[null]
5211female25.0101176555.4417E50C5[null]
5311male48.010PC 1757276.7292D33C3[null]
5411female35.0103697383.475C83SD[null]
5511male27.000PC 1757276.7292D49C3[null]
5611female24.0001176783.1583C54C7[null]
5711female52.0111274993.5B69S3[null]
5811female44.00111136157.9792B18C4[null]
5911female15.00124160211.3375B5S2[null]
6011male30.0101323657.75C78C11[null]
6111female31.01035273113.275D36C6[null]
6211female39.000PC 17758108.9C105C8[null]
6311female22.00111350961.9792B36C5[null]
6411male52.00011378630.5C104S6[null]
6511female43.00124160211.3375B3S2[null]
6611female33.00011015286.5B77S8[null]
6711male45.01116966134.5E34C3[null]
6811female40.01116966134.5E34C3[null]
6911male48.0101999652.0C126S5 7[null]
7011female[null]00PC 1758579.2[null]CD[null]
7111female35.000PC 17755512.3292[null]C3[null]
7211female60.01011081375.25D37C5[null]
7311male21.001PC 1759761.3792[null]CA[null]
7420male23.000C.A. 3103010.5[null]S[null][null]
7520male28.00024435826.0[null]S[null][null]
7620male60.0112975039.0[null]S[null][null]
7720female44.01024425226.0[null]S[null][null]
7820male29.010200326.0[null]S[null][null]
7920male18.000S.O.C. 1487973.5[null]S[null][null]
8020male18.000S.O.C. 1487973.5[null]S[null][null]
8120male54.0002840326.0[null]S[null][null]
8220male18.00023617113.0[null]S[null][null]
8320male36.00022923613.0[null]S[null]236
8420male34.0102866421.0[null]S[null][null]
8520male21.0102813311.5[null]S[null][null]
8620male21.0102813411.5[null]S[null][null]
8720male24.00023386613.0[null]S[null]155
8820male34.0001223313.0[null]S[null][null]
8920male30.00025065313.0[null]S[null]75
9020male44.00024874613.0[null]S[null]35
9120male49.01222084565.0[null]S[null][null]
9220male21.020S.O.C. 1487973.5[null]S[null][null]
9320male21.000S.O.C. 1487973.5[null]S[null][null]
9420female60.0102406526.0[null]S[null][null]
9520male24.020C.A. 3102931.5[null]S[null][null]
9620male22.020C.A. 3102931.5[null]S[null][null]
9720male35.00023373412.35[null]Q[null][null]
9820male31.000C.A. 1872310.5[null]S[null]165
9920male36.000SC/Paris 216312.875DC[null][null]
10020male[null]00SC/A.3 286115.5792[null]C[null][null]
Rows: 1-100 | Columns: 14

Note

VerticaPy offers a wide range of sample datasets that are ideal for training and testing purposes. You can explore the full list of available datasets in the Datasets, which provides detailed information on each dataset and how to use them effectively. These datasets are invaluable resources for honing your data analysis and machine learning skills within the VerticaPy environment.

Lets create a non-binary tree using vDataFrame.chaid() method.

tree = data.chaid("survived", ["sex", "fare"]).tree_

Our non-binary tree is ready, we will now provide information about classes and create a NonBinaryTree model.

classes = ["a", "b"]

model_nbt = NonBinaryTree(tree, classes)

Create a dataset.

data = [["male", 100], ["female", 20], ["female", 50]]

Making In-Memory Predictions

Use predict() method to do predictions.

model_nbt.predict(data)
Out[6]: array(['a', 'b', 'b'], dtype='<U1')

Use predict_proba() method to compute the predicted probabilities for each class.

model_nbt.predict_proba(data)
Out[7]: 
array([[0.82129278, 0.17870722],
       [0.3042328 , 0.6957672 ],
       [0.3042328 , 0.6957672 ]])

Deploy SQL Code

Let’s use the following column names:

cnames = ["sex", "fare"]

Use predict_sql() method to get the SQL code needed to deploy the model using its attributes.

model_nbt.predict_sql(cnames)
Out[9]: "(CASE WHEN sex = 'female' THEN (CASE WHEN fare <= 127.6 THEN 'b' WHEN fare <= 255.2 THEN 'b' WHEN fare <= 382.8 THEN 'b' WHEN fare <= 638.0 THEN 'b' ELSE NULL END) WHEN sex = 'male' THEN (CASE WHEN fare <= 129.36 THEN 'a' WHEN fare <= 258.72 THEN 'a' WHEN fare <= 388.08 THEN 'a' WHEN fare <= 517.44 THEN 'b' ELSE NULL END) ELSE NULL END)"

Use predict_proba_sql() method to get the SQL code needed to deploy the model that computes predicted probabilities.

model_nbt.predict_proba_sql(cnames)
Out[10]: 
["(CASE WHEN sex = 'female' THEN (CASE WHEN fare <= 127.6 THEN 0.304232804232804 WHEN fare <= 255.2 THEN 0.09375 WHEN fare <= 382.8 THEN 0.0 WHEN fare <= 638.0 THEN 0.0 ELSE NULL END) WHEN sex = 'male' THEN (CASE WHEN fare <= 129.36 THEN 0.821292775665399 WHEN fare <= 258.72 THEN 0.777777777777778 WHEN fare <= 388.08 THEN 0.75 WHEN fare <= 517.44 THEN 0.0 ELSE NULL END) ELSE NULL END)",
 "(CASE WHEN sex = 'female' THEN (CASE WHEN fare <= 127.6 THEN 0.695767195767196 WHEN fare <= 255.2 THEN 0.90625 WHEN fare <= 382.8 THEN 1.0 WHEN fare <= 638.0 THEN 1.0 ELSE NULL END) WHEN sex = 'male' THEN (CASE WHEN fare <= 129.36 THEN 0.178707224334601 WHEN fare <= 258.72 THEN 0.222222222222222 WHEN fare <= 388.08 THEN 0.25 WHEN fare <= 517.44 THEN 1.0 ELSE NULL END) ELSE NULL END)"]

Hint

This object can be pickled and used in any in-memory environment, just like SKLEARN models.

Drawing Tree

Use to_graphviz() method to generate code for a Graphviz tree.

model_nbt.to_graphviz()
Out[11]: 'digraph Tree {\ngraph [bgcolor="#FFFFFFDD"];\n0 [label="\\"sex\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n0 -> 1[label="= female", color="#000000", fontcolor="#000000"]\n1 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n1 -> 2[label="<= 127.6", color="#000000", fontcolor="#000000"]2 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#efc5b5"><b> prediction: b </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.3 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 0.7 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n1 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n1 -> 3[label="<= 255.2", color="#000000", fontcolor="#000000"]3 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#efc5b5"><b> prediction: b </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.09 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 0.91 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n1 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n1 -> 4[label="<= 382.8", color="#000000", fontcolor="#000000"]4 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#efc5b5"><b> prediction: b </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.0 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 1.0 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n1 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n1 -> 5[label="<= 638.0", color="#000000", fontcolor="#000000"]5 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#efc5b5"><b> prediction: b </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.0 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 1.0 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n0 [label="\\"sex\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n0 -> 6[label="= male", color="#000000", fontcolor="#000000"]\n6 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n6 -> 7[label="<= 129.36", color="#000000", fontcolor="#000000"]7 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#87cefa"><b> prediction: a </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.82 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 0.18 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n6 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n6 -> 8[label="<= 258.72", color="#000000", fontcolor="#000000"]8 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#87cefa"><b> prediction: a </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.78 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 0.22 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n6 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n6 -> 9[label="<= 388.08", color="#000000", fontcolor="#000000"]9 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#87cefa"><b> prediction: a </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.75 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 0.25 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n6 [label="\\"fare\\"", fillcolor="#FFFFFFDD", fontcolor="#000000", color="#000000"]\n6 -> 10[label="<= 517.44", color="#000000", fontcolor="#000000"]10 [label=<<table border="0" cellspacing="0"><tr><td port="port1" border="1" bgcolor="#efc5b5"><b> prediction: b </b></td></tr><tr><td port="port0" border="1" align="left"> prob(a): 0.0 </td></tr><tr><td port="port1" border="1" align="left"> prob(b): 1.0 </td></tr></table>>, fillcolor="#FFFFFFDD", fontcolor="#000000", shape="none", color="#000000"]\n}'

Use plot_tree() method to draw the input tree.

model_nbt.plot_tree()
../_images/machine_learning_memmodel_tree_NonBinaryTree.png

Important

plot_tree() requires the Graphviz module.

Note

The above example is a very basic one. For other more detailed examples and customization options, please see :ref:`chart_gallery.tree`_

__init__(tree: dict, classes: Annotated[list | ndarray, 'Array Like Structure'] | None = None) → None

Methods

__init__(tree[, classes])

get_attributes()

Returns the model attributes.

plot_tree([pic_path])

Draws the input tree.

predict(X)

Predicts using the CHAID model.

predict_proba(X)

Returns probabilities using the CHAID model.

predict_proba_sql(X)

Returns the SQL code needed to deploy the model probabilities.

predict_sql(X)

Returns the SQL code needed to deploy the model using its attributes.

set_attributes(**kwargs)

Sets the model attributes.

to_graphviz([classes_color, round_pred, ...])

Returns the code for a Graphviz tree.

Attributes

object_type

Must be overridden in child class