1515 < script src ="../../../_static/jquery.js?v=5d32c60e "> </ script >
1616 < script src ="../../../_static/_sphinx_javascript_frameworks_compat.js?v=2cd50e6c "> </ script >
1717 < script src ="../../../_static/documentation_options.js?v=5929fcd5 "> </ script >
18- < script src ="../../../_static/doctools.js?v=888ff710 "> </ script >
18+ < script src ="../../../_static/doctools.js?v=9a2dae69 "> </ script >
1919 < script src ="../../../_static/sphinx_highlight.js?v=dc90522c "> </ script >
2020 < script src ="../../../_static/js/theme.js "> </ script >
2121 < link rel ="index " title ="Index " href ="../../../genindex.html " />
7979 < div itemprop ="articleBody ">
8080
8181 < h1 > Source code for multimodal_transformers.data.tabular_torch_dataset</ h1 > < div class ="highlight "> < pre >
82- < span > </ span > < span class ="kn "> import</ span > < span class ="nn "> numpy</ span > < span class ="k "> as</ span > < span class ="nn "> np</ span >
82+ < span > </ span > < span class ="kn "> from</ span > < span class ="nn "> typing</ span > < span class ="kn "> import</ span > < span class ="n "> List</ span > < span class ="p "> ,</ span > < span class ="n "> Optional</ span > < span class ="p "> ,</ span > < span class ="n "> Union</ span >
83+
84+ < span class ="kn "> import</ span > < span class ="nn "> numpy</ span > < span class ="k "> as</ span > < span class ="nn "> np</ span >
85+ < span class ="kn "> import</ span > < span class ="nn "> pandas</ span > < span class ="k "> as</ span > < span class ="nn "> pd</ span >
8386< span class ="kn "> import</ span > < span class ="nn "> torch</ span >
87+ < span class ="kn "> import</ span > < span class ="nn "> transformers</ span >
8488< span class ="kn "> from</ span > < span class ="nn "> torch.utils.data</ span > < span class ="kn "> import</ span > < span class ="n "> Dataset</ span > < span class ="k "> as</ span > < span class ="n "> TorchDataset</ span >
8589
8690
@@ -91,31 +95,33 @@ <h1>Source code for multimodal_transformers.data.tabular_torch_dataset</h1><div
9195< span class ="sd "> :obj:`TorchDataset` wrapper for text dataset with categorical features</ span >
9296< span class ="sd "> and numerical features</ span >
9397
94- < span class ="sd "> Parameters :</ span >
95- < span class ="sd "> encodings (:class:` transformers.BatchEncoding`): </ span >
96- < span class =" sd " > The output from encode_plus() and batch_encode() methods (tokens, attention_masks, etc) of </ span >
97- < span class ="sd "> a transformers.PreTrainedTokenizer </ span >
98- < span class ="sd "> categorical_feats (:class:`numpy.ndarray`, of shape :obj: `(n_examples, categorical feat dim)`, `optional`, defaults to :obj:`None`): </ span >
99- < span class =" sd " > An array containing the preprocessed categorical features </ span >
100- < span class ="sd "> numerical_feats (:class:`numpy.ndarray`, of shape :obj:`(n_examples, numerical feat dim)`, `optional`, defaults to :obj:`None`) :</ span >
101- < span class ="sd "> An array containing the preprocessed numerical features</ span >
102- < span class =" sd " > labels (:class: list` or `numpy.ndarray`, `optional`, defaults to :obj:`None`): </ span >
103- < span class ="sd "> The labels of the training examples </ span >
104- < span class ="sd "> df (:class:`pandas.DataFrame`, `optional`, defaults to :obj:`None`): </ span >
105- < span class =" sd " > Model configuration class with all the parameters of the model. </ span >
106- < span class ="sd "> This object must also have a tabular_config member variable that is a </ span >
107- < span class ="sd "> TabularConfig instance specifying the configs for TabularFeatCombiner </ span >
98+ < span class ="sd "> :param encodings :</ span >
99+ < span class ="sd "> The output from `encode_plus()` and `batch_encode()` methods (tokens, attention_masks, etc.) of a ` transformers.PreTrainedTokenizer`. </ span >
100+
101+ < span class ="sd "> :param categorical_feats: </ span >
102+ < span class ="sd "> An array containing the preprocessed categorical features. Shape: `(n_examples, categorical feat dim)`. </ span >
103+
104+ < span class ="sd "> :param numerical_feats :</ span >
105+ < span class ="sd "> An array containing the preprocessed numerical features. Shape: `(n_examples, numerical feat dim)`. </ span >
106+
107+ < span class ="sd "> :param labels: </ span >
108+ < span class ="sd "> The labels of the training examples. </ span >
109+
110+ < span class ="sd "> :param df: </ span >
111+ < span class ="sd "> The original dataset. Optional and used only to save the original dataset with the preprocessed dataset. </ span >
108112
113+ < span class ="sd "> :param label_list:</ span >
114+ < span class ="sd "> A list of class names for each unique class in labels.</ span >
109115< span class ="sd "> """</ span >
110116
111117 < span class ="k "> def</ span > < span class ="fm "> __init__</ span > < span class ="p "> (</ span >
112118 < span class ="bp "> self</ span > < span class ="p "> ,</ span >
113- < span class ="n "> encodings</ span > < span class ="p "> ,</ span >
114- < span class ="n "> categorical_feats</ span > < span class ="p "> ,</ span >
115- < span class ="n "> numerical_feats</ span > < span class ="p "> ,</ span >
116- < span class ="n "> labels</ span > < span class ="o "> = </ span > < span class ="kc "> None</ span > < span class ="p "> ,</ span >
117- < span class ="n "> df</ span > < span class ="o "> = </ span > < span class ="kc "> None</ span > < span class ="p "> ,</ span >
118- < span class ="n "> label_list</ span > < span class ="o "> =</ span > < span class ="kc "> None</ span > < span class ="p "> ,</ span >
119+ < span class ="n "> encodings</ span > < span class ="p "> : </ span > < span class =" n " > transformers </ span > < span class =" o " > . </ span > < span class =" n " > BatchEncoding </ span > < span class =" p " > ,</ span >
120+ < span class ="n "> categorical_feats</ span > < span class ="p "> : </ span > < span class =" n " > Optional </ span > < span class =" p " > [ </ span > < span class =" n " > pd </ span > < span class =" o " > . </ span > < span class =" n " > DataFrame </ span > < span class =" p " > ] ,</ span >
121+ < span class ="n "> numerical_feats</ span > < span class ="p "> : </ span > < span class =" n " > Optional </ span > < span class =" p " > [ </ span > < span class =" n " > pd </ span > < span class =" o " > . </ span > < span class =" n " > DataFrame </ span > < span class =" p " > ] ,</ span >
122+ < span class ="n "> labels</ span > < span class ="p " > : </ span > < span class =" n " > Optional </ span > < span class =" p " > [ </ span > < span class =" n " > Union </ span > < span class =" p " > [ </ span > < span class =" n " > List </ span > < span class =" p " > , </ span > < span class =" n " > np </ span > < span class =" o "> . </ span > < span class =" n " > ndarray </ span > < span class =" p " > ]] </ span > < span class =" o " > = </ span > < span class ="kc "> None</ span > < span class ="p "> ,</ span >
123+ < span class ="n "> df</ span > < span class ="p " > : </ span > < span class =" n " > Optional </ span > < span class =" p " > [ </ span > < span class =" n " > pd </ span > < span class =" o "> . </ span > < span class =" n " > DataFrame </ span > < span class =" p " > ] </ span > < span class =" o " > = </ span > < span class ="kc "> None</ span > < span class ="p "> ,</ span >
124+ < span class ="n "> label_list</ span > < span class ="p " > : </ span > < span class =" n " > Optional </ span > < span class =" p " > [ </ span > < span class =" n " > List </ span > < span class =" p " > [ </ span > < span class =" n " > Union </ span > < span class =" p " > [ </ span > < span class =" nb " > str </ span > < span class =" p " > ]]] </ span > < span class =" o "> =</ span > < span class ="kc "> None</ span > < span class ="p "> ,</ span >
119125 < span class ="p "> ):</ span >
120126 < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> df</ span > < span class ="o "> =</ span > < span class ="n "> df</ span >
121127 < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> encodings</ span > < span class ="o "> =</ span > < span class ="n "> encodings</ span >
@@ -128,13 +134,13 @@ <h1>Source code for multimodal_transformers.data.tabular_torch_dataset</h1><div
128134 < span class ="k "> else</ span > < span class ="p "> [</ span > < span class ="n "> i</ span > < span class ="k "> for</ span > < span class ="n "> i</ span > < span class ="ow "> in</ span > < span class ="nb "> range</ span > < span class ="p "> (</ span > < span class ="nb "> len</ span > < span class ="p "> (</ span > < span class ="n "> np</ span > < span class ="o "> .</ span > < span class ="n "> unique</ span > < span class ="p "> (</ span > < span class ="n "> labels</ span > < span class ="p "> )))]</ span >
129135 < span class ="p "> )</ span >
130136
131- < span class ="k "> def</ span > < span class ="fm "> __getitem__</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="p "> ,</ span > < span class ="n "> idx</ span > < span class ="p "> ):</ span >
137+ < span class ="k "> def</ span > < span class ="fm "> __getitem__</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="p "> ,</ span > < span class ="n "> idx</ span > < span class ="p "> : </ span > < span class =" nb " > int </ span > < span class =" p " > ):</ span >
132138 < span class ="n "> item</ span > < span class ="o "> =</ span > < span class ="p "> {</ span > < span class ="n "> key</ span > < span class ="p "> :</ span > < span class ="n "> torch</ span > < span class ="o "> .</ span > < span class ="n "> tensor</ span > < span class ="p "> (</ span > < span class ="n "> val</ span > < span class ="p "> [</ span > < span class ="n "> idx</ span > < span class ="p "> ])</ span > < span class ="k "> for</ span > < span class ="n "> key</ span > < span class ="p "> ,</ span > < span class ="n "> val</ span > < span class ="ow "> in</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> encodings</ span > < span class ="o "> .</ span > < span class ="n "> items</ span > < span class ="p "> ()}</ span >
133139 < span class ="n "> item</ span > < span class ="p "> [</ span > < span class ="s2 "> "labels"</ span > < span class ="p "> ]</ span > < span class ="o "> =</ span > < span class ="p "> (</ span >
134140 < span class ="n "> torch</ span > < span class ="o "> .</ span > < span class ="n "> tensor</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> labels</ span > < span class ="p "> [</ span > < span class ="n "> idx</ span > < span class ="p "> ])</ span > < span class ="k "> if</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> labels</ span > < span class ="ow "> is</ span > < span class ="ow "> not</ span > < span class ="kc "> None</ span > < span class ="k "> else</ span > < span class ="kc "> None</ span >
135141 < span class ="p "> )</ span >
136142 < span class ="n "> item</ span > < span class ="p "> [</ span > < span class ="s2 "> "cat_feats"</ span > < span class ="p "> ]</ span > < span class ="o "> =</ span > < span class ="p "> (</ span >
137- < span class ="n "> torch</ span > < span class ="o "> .</ span > < span class ="n "> tensor</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> cat_feats</ span > < span class ="p "> [</ span > < span class ="n "> idx</ span > < span class ="p "> ])</ span > < span class ="o "> .</ span > < span class ="n "> float</ span > < span class ="p "> ()</ span >
143+ < span class ="n "> torch</ span > < span class ="o "> .</ span > < span class ="n "> tensor</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> cat_feats</ span > < span class ="o " > . </ span > < span class =" n " > iloc </ span > < span class =" p "> [</ span > < span class ="n "> idx</ span > < span class ="p "> ])</ span > < span class ="o "> .</ span > < span class ="n "> float</ span > < span class ="p "> ()</ span >
138144 < span class ="k "> if</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> cat_feats</ span > < span class ="ow "> is</ span > < span class ="ow "> not</ span > < span class ="kc "> None</ span >
139145 < span class ="k "> else</ span > < span class ="n "> torch</ span > < span class ="o "> .</ span > < span class ="n "> zeros</ span > < span class ="p "> (</ span > < span class ="mi "> 0</ span > < span class ="p "> )</ span >
140146 < span class ="p "> )</ span >
@@ -145,13 +151,13 @@ <h1>Source code for multimodal_transformers.data.tabular_torch_dataset</h1><div
145151 < span class ="p "> )</ span >
146152 < span class ="k "> return</ span > < span class ="n "> item</ span >
147153
148- < span class ="k "> def</ span > < span class ="fm "> __len__</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="p "> ):</ span >
154+ < span class ="k "> def</ span > < span class ="fm "> __len__</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="p "> )</ span > < span class =" o " > -> </ span > < span class =" nb " > int </ span > < span class =" p " > :</ span >
149155 < span class ="k "> return</ span > < span class ="nb "> len</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> encodings</ span > < span class ="p "> [</ span > < span class ="s2 "> "input_ids"</ span > < span class ="p "> ])</ span >
150156
151157< div class ="viewcode-block " id ="TorchTabularTextDataset.get_labels ">
152158< a class ="viewcode-back " href ="../../../modules/data.html#multimodal_transformers.data.TorchTabularTextDataset.get_labels "> [docs]</ a >
153- < span class ="k "> def</ span > < span class ="nf "> get_labels</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="p "> ):</ span >
154- < span class ="w "> </ span > < span class ="sd "> """returns the label names for classification"""</ span >
159+ < span class ="k "> def</ span > < span class ="nf "> get_labels</ span > < span class ="p "> (</ span > < span class ="bp "> self</ span > < span class ="p "> )</ span > < span class =" o " > -> </ span > < span class =" n " > Optional </ span > < span class =" p " > [ </ span > < span class =" n " > List </ span > < span class =" p " > [ </ span > < span class =" n " > Union </ span > < span class =" p " > [ </ span > < span class =" nb " > str </ span > < span class =" p " > ]]] :</ span >
160+ < span class ="w "> </ span > < span class ="sd "> """Returns the label names for classification. """</ span >
155161 < span class ="k "> return</ span > < span class ="bp "> self</ span > < span class ="o "> .</ span > < span class ="n "> label_list</ span > </ div >
156162</ div >
157163
0 commit comments