CoolFace
Modelpublic

Penghaoo/workspace

sourceHugging Facegemmaupdated 2y agoView on Hugging Face
0likes5downloads
dataloader.py297 linesDownload Raw Back to root
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