From dd74a6d2820c6689263f032a47d9e241eef5d605 Mon Sep 17 00:00:00 2001 From: "NAUSICAA\\Julian" Date: Sun, 24 Feb 2019 23:15:20 -0300 Subject: [PATCH] Remove previous multi field patch --- ulmfit/train_clas.py | 14 -------------- 1 file changed, 14 deletions(-) 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