diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index bc04885..1dfb2ce 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -148,16 +148,6 @@ class CLSHyperParams(LMHyperParams): **kwargs) return self.databunches(bs, **data) - def merge_cols(self, df): - if len(df.columns) <= 2: - return df - ndf = df[[0,1]].copy() - for i in range(2, len(df.columns)): - ndf[1] += ("\n" + FLD + "\n") + df[i].fillna(" ") - - assert ndf[1].isna().sum().sum() == 0, f"You have NaN values in column(s) of your dataframe, please fix it." - return ndf - def load_data(self, lang='', **kwargs): prefix = '' if lang == '' else lang+'.' trn_df = pd.read_csv(self.dataset_path / f'{prefix}train.csv', header=None) @@ -176,10 +166,6 @@ class CLSHyperParams(LMHyperParams): val_len = max(int(len(trn_df) * 0.1), 2) trn_len = len(trn_df) - val_len trn_df, val_df = trn_df[:trn_len], trn_df[trn_len:] - trn_df = self.merge_cols(trn_df) - val_df = self.merge_cols(val_df) - tst_df = self.merge_cols(tst_df) - unsup_df = self.merge_cols(unsup_df) kwargs.update(dict(trn_df=trn_df, val_df=val_df, tst_df=tst_df, unsup_df=unsup_df)) return kwargs