Penghaoo/workspace
05
1from datasets import Dataset, DatasetDict2import pandas as pd3import numpy as np4import glob5from sklearn.model_selection import train_test_split6import re7 8datapath = '/cluster/work/lawecon/Work/penghao/dataset/stories/'9pairpath = '../../../work/lawecon/Work/penghao/pairs.csv'10#3600 ->time lags11 12 13class StoryPairDataset(Dataset):14 def __init__(self, datapath, pairpath, tokenizer, task, used_dataset_size=-1, train_test_split=0.1,15 split_by='random',16 max_len=4096*2, mode='m3', max_time_window=3000, least_likes=5, margin=True):17 self.datapath = datapath18 print(self.datapath)19 self.train_test_split = train_test_split20 self.pairpath = pairpath21 self.tokenizer = tokenizer22 self.max_len = max_len23 self.split_by = split_by24 self.least_likes = least_likes25 self.max_time_window = max_time_window26 self.used_dataset_size = used_dataset_size27 if mode == 'm2':28 self.max_time_window = 1200960029 else:30 self.max_time_window = max_time_window31 self.pair = self.load_pair()32 33 self.task = task34 self.margin = margin35 self.stories = self.load_stories(self.datapath)36 print(self.stories.columns)37 print(len(self.stories))38 39 40 # turn df into dataset41 42 # self.dataset = datasets.Dataset.from_pandas(self.df)43 self.train, self.test = self.train_test_split__()44 self.train = self.marginInclude(self.train)45 self.test = self.marginInclude(self.test)46 # combine train and test to a single dataset, before train and test47 self.dataset = self.make_dataset()48 print('current setting mode is ', mode)49 print('currnet setting split_by is ', split_by)50 print('current setting least_likes is ', least_likes)51 52 53 54 def load_stories(self, path):55 stories = pd.DataFrame()56 #print(f"Reading stories from {path}...")57 for file in glob.glob(path + '*.csv'):58 #print(f"Reading {file}...")59 try:60 # Read the CSV file into a DataFrame61 df = pd.read_csv(file)62 63 # Check if the DataFrame is empty or not64 if df.empty:65 print(f"Warning: {file} is empty or not readable.")66 continue67 # Concatenate the DataFrames68 stories = pd.concat([stories, df], ignore_index=True)69 except pd.errors.EmptyDataError:70 # print(f"Error: {file} is empty or not readable.")71 pass72 except pd.errors.ParserError:73 print(f"Error: {file} cannot be parsed.")74 except Exception as e:75 print(f"Error: An unexpected error occurred while processing {file}. Details: {str(e)}")76 # contain Index(['prompt_id', 'prompt', 'story_id', 'story_title', 'story_author', 'story_url', 'link', 'genre', 'is_sensitive', 'categories', 'likes', 'story_text', 'posted_date', 'comments'], dtype='object')77 78 return stories79 80 def load_pair(self):81 82 pair = pd.read_csv(self.pairpath)83 # contain the colums of prompt_id, story1_id, story2_id, rel, time_lag, least_likes84 85 pair = pair[pair['time_lag'] <= self.max_time_window]86 print('the max of tima lag is ', pair['time_lag'].max())87 88 pair = pair[pair['least_likes'] >= self.least_likes]89 # swap the order of story1 and story2 if rel is negative, and makes rel positive90 pair.loc[pair['rel'] < 0, ['story1_id', 'story2_id']] = pair.loc[91 pair['rel'] < 0, ['story2_id', 'story1_id']].values92 pair['rel'] = abs(pair['rel'])93 # filter the pair if they have same story id 94 pair = pair[pair['story1_id'] != pair['story2_id']]95 if self.used_dataset_size == -1:96 self.used_dataset_size = len(pair)97 else:98 pair = pair.sample(n=self.used_dataset_size)99 print('the total number of pairs is ', len(pair))100 # remove the duplicate pairs 101 pair = pair.drop_duplicates(subset=['story1_id', 'story2_id'])102 #remove the rel = 0103 pair = pair[pair['rel'] != 0]104 print('the number of effective pairs is ', len(pair))105 return pair106 107 def marginInclude(self, df):108 if self.margin:109 # drop the column of rel110 df = df.drop(columns=['rel'])111 else:112 # rename rel to margin113 df = df.rename(columns={'rel': 'margin'})114 return df115 116 def train_test_split__(self):117 '''118 split the pairs into train and test set119 :return:120 '''121 test_size = round(len(self.pair) * self.train_test_split)122 123 if self.split_by == 'time':124 # give the pair the information of year according to the story_id125 self.stories['posted_date'] = pd.to_datetime(self.stories['posted_date'])126 #convert datetime64[ns] to comparable format, e.g. 2021-04-27 23:29:00 -> 20210427127 self.stories['posted_date'] = self.stories['posted_date'].dt.strftime('%Y%m%d')128 # the time after 2022 is test set129 130 131 test = self.pair[self.pair['story1_id'].apply(lambda x: int(self.stories[self.stories['story_id'] == x]['posted_date'].values[0]) > 20220000)]132 train = self.pair[self.pair['story1_id'].apply(lambda x: int(self.stories[self.stories['story_id'] == x]['posted_date'].values[0]) <= 20220000)]133 print('the number of test set is ', len(test))134 print('the number of train set is ', len(train))135 print('the ratio of test set is ', len(test) / (len(test) + len(train)))136 137 elif self.split_by == 'random':138 139 train, test = train_test_split(self.pair, test_size=self.train_test_split)140 141 # covert to huggingface dataset142 143 144 elif self.split_by == 'genre':145 146 # count the number of pairs for each category147 # give the pair the information of category according to the story_id148 self.pair['genre'] = self.pair['story1_id'].apply(149 lambda x: self.stories[self.stories['story_id'] == x]['genre'].values[0])150 genre = {}151 for c in self.pair['genre'].unique():152 genre[c] = len(self.pair[self.pair['genre'] == c])153 # select the category to nearest to 10 per cent of the total154 genre = dict(sorted(genre.items(), key=lambda item: item[1], reverse=True))#sort the genre by the number of pairs from high to low155 print(genre)156 total = sum(genre.values())157 #select the close genre to 10% of the total158 test_genre = []159 test_count = 0160 while test_count < total * self.train_test_split:161 test_genre.append(list(genre.keys())[0])162 test_count += genre[list(genre.keys())[0]]163 del genre[list(genre.keys())[0]]164 if test_count + genre[list(genre.keys())[0]] > total * self.train_test_split:165 break166 167 test = self.pair[self.pair['genre'].apply(lambda x: x in test_genre)]168 train = self.pair[self.pair['genre'].apply(lambda x: x not in test_genre)]169 print('the genre of test set is ', test_genre)170 print('the percentage of test set is ', test_count / total,'where total is ', total)171 172 elif self.split_by == 'chaos':173 #instead using the pairs, we randomly assign the story id to replace the old story id from that prompt174 for i in range(len(self.pair)):175 self.pair.at[i, 'story1_id'] = np.random.choice(self.stories[self.stories['prompt_id'] == self.pair.at[i, 'prompt_id']]['story_id'].values)176 self.pair.at[i, 'story2_id'] = np.random.choice(self.stories[self.stories['prompt_id'] == self.pair.at[i, 'prompt_id']]['story_id'].values)177 train, test = train_test_split(self.pair, test_size=self.train_test_split)178 return train, test179 180 def apply_template_to_text(self, row):181 182 # Ensure proper access to columns in pair183 prompt_id, story1_id, story2_id = row[['prompt_id', 'story1_id', 'story2_id']]184 185 # Extract text based on IDs186 187 chosen_prompt = self.stories[self.stories['prompt_id'] == prompt_id]['prompt']188 chosen_prompt = chosen_prompt.values[0]189 chosen_story = self.stories[self.stories['story_id'] == story1_id]['story_title'].values[0] + '/n' + \190 self.stories[self.stories['story_id'] == story1_id]['story_text'].values[0]191 192 rejected_prompt = self.stories[self.stories['prompt_id'] == prompt_id]['prompt']193 rejected_prompt = rejected_prompt.values[0]194 rejected_story = self.stories[self.stories['story_id'] == story2_id]['story_title'].values[0] + '/n' + \195 self.stories[self.stories['story_id'] == story2_id]['story_text'].values[0]196 197 # Create chosen and rejected text dictionaries198 chosen_text = [{'role': 'user', 'content': chosen_prompt},199 {'role': 'assistant', 'content': chosen_story}]200 201 rejected_text = [{'role': 'user', 'content': rejected_prompt},202 {'role': 'assistant', 'content': rejected_story}]203 204 # Apply tokenizer to chosen and rejected text205 chosen_text = self.tokenizer.apply_chat_template(chosen_text, tokenize=False)206 rejected_text = self.tokenizer.apply_chat_template(rejected_text, tokenize=False)207 208 res = {}209 res['chosen_text'] = chosen_text210 res['rejected_text'] = rejected_text211 #add eos and bos token212 res['chosen_text'] = self.tokenizer.bos_token + res['chosen_text'] + self.tokenizer.eos_token213 res['rejected_text'] = self.tokenizer.bos_token + res['rejected_text'] + self.tokenizer.eos_token214 215 res['text'] = chosen_text216 #add eos and bos token217 res['text'] = self.tokenizer.bos_token + res['text'] + self.tokenizer.eos_token218 if 'gemma' in self.tokenizer.name_or_path:219 split_words = '<|im_start|>assistant\n'220 elif 'mistral' in self.tokenizer.name_or_path or 'llama' in self.tokenizer.name_or_path:221 split_words = '[/INST]'222 223 chosen_text_tmp = chosen_text.split(split_words)[-1]224 prompt_text = chosen_text.replace(chosen_text_tmp, '')225 chosen_text = chosen_text_tmp226 227 rejected_text = rejected_text.split(split_words)[-1]228 res['prompt'] = prompt_text229 res['chosen'] = chosen_text230 res['rejected'] = rejected_text231 # add bos and eos token232 res['prompt'] = self.tokenizer.bos_token + res['prompt']233 res['chosen'] = res['chosen'] + self.tokenizer.eos_token234 res['rejected'] = res['rejected'] + self.tokenizer.eos_token235 return res236 237 def convert_sft(self,df):238 #collect all the story id in the pair239 story_ids = list(set(df['story1_id'].values) | set(df['story2_id'].values))240 #now make new train and test set as story_ids as story1_id and story2_id241 df = pd.DataFrame()242 df['story1_id'] = story_ids243 df['story2_id'] = df['story1_id']244 #reload stories245 #self.stories = self.load_stories(self.datapath)246 # get prompt_id from the pair247 def get_prompt_id(x):248 return self.stories[self.stories['story_id'] == x]['prompt_id'].values[0]249 df['prompt_id'] = df['story1_id'].apply(lambda x: get_prompt_id(x))250 return df251 252 253 254 def make_dataset(self):255 # reset the index256 self.train.reset_index(drop=True, inplace=True)257 self.test.reset_index(drop=True, inplace=True)258 entries = []259 if self.task == 'rm':260 entries = ['chosen_text', 'rejected_text']261 elif self.task == 'dpo':262 entries = ['prompt', 'chosen', 'rejected']263 elif self.task == 'sft':264 self.train = self.convert_sft(self.train)265 self.test = self.convert_sft(self.test)266 entries = ['text']267 268 print('the columns of train is ', self.train.columns)269 for index, row in self.train.iterrows():270 res = self.apply_template_to_text(row)271 for e in entries:272 self.train.at[index, e] = res[e]273 274 for index, row in self.test.iterrows():275 res = self.apply_template_to_text(row)276 for e in entries:277 self.test.at[index, e] = res[e]278 279 print('the first example of train is ', self.train.iloc[0])280 #since the we aggred on max_len = 8192, we need to filter this281 282 if self.margin:283 entries.append('margin')284 285 train_dataset = Dataset.from_pandas(self.train[entries])286 test_dataset = Dataset.from_pandas(self.test[entries])287 288 return DatasetDict({'train': train_dataset, 'test': test_dataset})289 290 def save_dataset(self, path):291 '''292 save the dataset to the readsy folder293 :param path:294 :return:295 '''296 self.dataset.save_to_disk('../' + path)297 