From 7dc7aac327e69c1b4c5372240032dc4846cefaf7 Mon Sep 17 00:00:00 2001 From: Piotr Czapla Date: Mon, 18 Feb 2019 21:50:04 +0100 Subject: [PATCH] Merge all columns in classification task into first column This should fix CLS issues. --- ulmfit/train_clas.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/ulmfit/train_clas.py b/ulmfit/train_clas.py index 12307be..b06bbeb 100644 --- a/ulmfit/train_clas.py +++ b/ulmfit/train_clas.py @@ -143,6 +143,16 @@ 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) @@ -161,7 +171,10 @@ 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