cleanrl/summarize_from_feedback_oai_preprocessing_1704321749
Dataset Card for "summarize_from_feedback_oai_preprocessing_1704321749" More Information needed
169
1import multiprocessing2import os3import time4from dataclasses import dataclass, field5from pprint import pformat6from typing import Dict, Optional7 8import matplotlib.pyplot as plt9import pandas as pd10import tyro11from datasets import load_dataset12from huggingface_hub import HfApi13from huggingface_hub.repocard import RepoCard14from rich.pretty import pprint15from transformers import AutoTokenizer16 17api = HfApi()18 19 20"""21poetry run python lm_human_preference_details/tldr_dataset.py22poetry run python lm_human_preference_details/tldr_dataset.py \23 --base_model=EleutherAI/pythia-160m \24 --max_sft_response_length=53 \25 --max_sft_query_response_length=562 \26 --max-rm-response-length=169 \27 --max_rm_query_response_length=63828poetry run python lm_human_preference_details/tldr_dataset.py \29 --base_model=EleutherAI/pythia-160m \30 --max_sft_response_length=48 \31 --max_sft_query_response_length=560 \32 --max-rm-response-length=48 \33 --max_rm_query_response_length=56034 35poetry run python lm_human_preference_details/tldr_dataset.py \36 --base_model=EleutherAI/pythia-160m \37 --max_sft_response_length=53 \38 --max_sft_query_response_length=562 \39 --max-rm-response-length=169 \40 --max_rm_query_response_length=638 \41 --hf_entity=cleanrl \42 --push_to_hub \43 --oai_params.padding=""44poetry run python lm_human_preference_details/tldr_dataset.py \45 --base_model=EleutherAI/pythia-160m \46 --max_sft_response_length=48 \47 --max_sft_query_response_length=560 \48 --max-rm-response-length=48 \49 --max_rm_query_response_length=560 \50 --push_to_hub \51 --oai_params.padding=""52"""53 54 55@dataclass56class TaskQueryHParams:57 length: int = 51258 format_str: Optional[59 str60 ] = "SUBREDDIT: r/{subreddit}\n\nTITLE: {title}\n\nPOST: {post}\n\nTL;DR:" # if underlying dataset yields dicts, can format arbitrarily61 truncate_field: Optional[str] = "post"62 truncate_text: Optional[str] = "\n"63 padding: Optional[str] = " " # empty spaces64 pad_side: Optional[str] = "left"65 66 67@dataclass68class Args:69 base_model: str = "gpt2" # EleutherAI/pythia-160m70 max_sft_response_length: int = 48 # 5371 max_sft_query_response_length: int = 512 + 48 # 56572 max_rm_response_length: int = 153 # 16973 max_rm_query_response_length: int = 512 + 153 # 66574 hf_entity: str = None75 push_to_hub: bool = False76 oai_params: TaskQueryHParams = field(default_factory=TaskQueryHParams)77 78 79def _ensure_length(toks, l, pad_sequence=None, pad_side=None, truncate_side=None):80 assert pad_side in (None, "left", "right")81 assert truncate_side in (None, "left", "right")82 if len(toks) < l:83 assert pad_sequence is not None84 pad_amt = l - len(toks)85 assert len(pad_sequence) >= pad_amt, f"{len(pad_sequence)} < {pad_amt}"86 if pad_side is None:87 assert len(toks) == l, f"Needed to pad! {len(toks)} < {l}"88 return toks89 elif pad_side == "left":90 return pad_sequence[-pad_amt:] + toks91 else:92 assert pad_side == "right"93 return toks + pad_sequence[:pad_amt]94 if truncate_side is None:95 assert len(toks) == l, f"Needed to truncate! {len(toks)} > {l}"96 return toks97 elif truncate_side == "left":98 return toks[-l:]99 else:100 assert truncate_side == "right"101 return toks[:l]102 103 104def _get_query_padding_for_task(encoder, hparams: TaskQueryHParams):105 return hparams.padding * hparams.length106 107 108def process_query(query_info: Dict[str, str], *, encoder, hparams: TaskQueryHParams, pad_sequence=None):109 if pad_sequence is None:110 pad_sequence = _get_query_padding_for_task(encoder, hparams)111 if isinstance(query_info, str):112 query_info = dict(query=query_info)113 else:114 # copy to avoid mutating input115 query_info = dict(**query_info)116 117 format_str = hparams.format_str or "{query}"118 query_tokens = encoder.encode(format_str.format(**query_info))119 truncate_field = hparams.truncate_field or "query"120 121 if truncate_field not in query_info:122 raise ValueError(f"Could not truncate field {truncate_field}, found fields: {query_info.keys()}!")123 while len(query_tokens) > hparams.length:124 if not len(query_info[truncate_field]):125 raise ValueError("Could not truncate enough!")126 127 i = -1 # default to just remove one character128 if hparams.truncate_text:129 try:130 i = query_info[truncate_field].rindex(hparams.truncate_text)131 except ValueError:132 pass133 query_info[truncate_field] = query_info[truncate_field][:i]134 query_tokens = encoder.encode(format_str.format(**query_info))135 136 query_token = _ensure_length(query_tokens, hparams.length, pad_side=hparams.pad_side, pad_sequence=pad_sequence)137 query = encoder.decode(query_token, skip_special_tokens=True).lstrip()138 return dict(139 query_token=query_token,140 query=query,141 )142 143 144if __name__ == "__main__":145 args = tyro.cli(Args)146 if args.hf_entity is None:147 args.hf_entity = api.whoami()["name"]148 assert isinstance(args.hf_entity, str)149 tokenizer = AutoTokenizer.from_pretrained(args.base_model)150 tokenizer.add_special_tokens({"pad_token": "[PAD]"})151 if len(args.oai_params.padding) > 0:152 args.oai_params.padding = tokenizer.encode(args.oai_params.padding)153 else:154 args.oai_params.padding = [tokenizer.pad_token_id]155 pprint(args.oai_params)156 timestamp = int(time.time())157 sft_ds = load_dataset("vwxyzjn/summarize_from_feedback_tldr_3_filtered")158 159 def process_query_data(x):160 # the `x['summary']` in `vwxyzjn/summarize_from_feedback_tldr_3_filtered`161 # DOES NOT HAVE a leading space so we are adding the leading space and162 # `<|endoftext|>` token163 reference_response = f" {x['summary']}<|endoftext|>"164 y = {165 **process_query(x, encoder=tokenizer, hparams=args.oai_params),166 "reference_response": reference_response,167 "reference_response_token": tokenizer.encode(168 reference_response,169 padding="max_length",170 max_length=args.max_sft_response_length,171 truncation=True,172 ),173 "reference_response_token_len": len(tokenizer.encode(reference_response)),174 }175 y["query_reference_response"] = y["query"].strip() + y["reference_response"]176 y["query_reference_response_token"] = tokenizer.encode(177 y["query_reference_response"],178 padding="max_length",179 max_length=args.max_sft_query_response_length,180 truncation=True,181 )182 y["query_reference_response_token_len"] = len(tokenizer.encode(y["query_reference_response"]))183 return y184 185 sft_ds = sft_ds.map(process_query_data, load_from_cache_file=False, num_proc=multiprocessing.cpu_count())186 if args.push_to_hub:187 sft_ds.push_to_hub(f"{args.hf_entity}/summarize_from_feedback_tldr_3_filtered_oai_preprocessing_{timestamp}")188 sft_card = RepoCard.load(189 f"{args.hf_entity}/summarize_from_feedback_tldr_3_filtered_oai_preprocessing_{timestamp}",190 repo_type="dataset",191 )192 sft_card.text = f"""\193# TL;DR SFT Dataset for OpenAI's [Summarize from Feedback](https://openai.com/blog/summarization/) task194 195The dataset is directly taken from https://github.com/openai/summarize-from-feedback/tree/700967448d10004279f138666442bf1497d0e705#reddit-tldr-dataset196 197These columns are taken directly from the aforementioned dataset:198 199* **id**: unique identifier for the post200* **subreddit**: subreddit the post was taken from201* **title**: title of the post202* **post**: body of the post203* **summary**: summary of the post204* **reference_response**: reference response for the post205 206These columns are added by this preprocessing script:207* **query**: length-limited query for summarization: OAI pre-processes the main text (title + subreddit + post), ensuring it has only 512 tokens; if the main text is too long, then it tries to truncate at the last `\n`. If it's too short it pads the main text ([summarize_from_feedback/tasks.py#L98-L165](https://github.com/openai/summarize-from-feedback/blob/700967448d10004279f138666442bf1497d0e705/summarize_from_feedback/tasks.py#L98-L165)). Padding is either space or `[PAD]` token (see Args below).208* **query_token**: tokenized version of `query`209* **reference_response_token**: tokenized version of `reference_response`210* **reference_response_token_len**: length of `reference_response_token`211* **query_reference_response**: concatenation of `query.strip()` and `reference_response`212* **query_reference_response_token**: tokenized version of `query_reference_response`, up to `max_sft_query_response_length` tokens213* **query_reference_response_token_len**: length of `query_reference_response_token`214 215 216# Args217 218```python219{pformat(vars(args))}220{pformat(vars(args.oai_params))}221```222"""223 sft_card.push_to_hub(224 f"{args.hf_entity}/summarize_from_feedback_tldr_3_filtered_oai_preprocessing_{timestamp}",225 repo_type="dataset",226 )227 228 label_ds = load_dataset("openai/summarize_from_feedback", "comparisons")229 230 def process_response_data(x):231 # the `x['summaries'][0]['text']` in `openai/summarize_from_feedback` `comaprisons`232 # DOES HAVE a leading space so we are just adding the `<|endoftext|>` token233 response0 = f"{x['summaries'][0]['text']}<|endoftext|>"234 response1 = f"{x['summaries'][1]['text']}<|endoftext|>"235 response0_policy = x["summaries"][0]["policy"]236 response1_policy = x["summaries"][1]["policy"]237 policies = "--".join(sorted([response0_policy, response1_policy]))238 y = {239 **process_query(x["info"], encoder=tokenizer, hparams=args.oai_params),240 "response0": response0,241 "response0_token": tokenizer.encode(242 response0, padding="max_length", max_length=args.max_rm_response_length, truncation=True243 ),244 "response0_token_len": len(tokenizer.encode(response0)),245 "response1": response1,246 "response1_token": tokenizer.encode(247 response1, padding="max_length", max_length=args.max_rm_response_length, truncation=True248 ),249 "response1_token_len": len(tokenizer.encode(response1)),250 "response0_policy": response0_policy,251 "response1_policy": response1_policy,252 "policies": policies,253 }254 y["query_response0"] = y["query"].strip() + y["response0"]255 y["query_response0_token"] = tokenizer.encode(256 y["query_response0"], padding="max_length", max_length=args.max_rm_query_response_length, truncation=True257 )258 y["query_response0_token_len"] = len(tokenizer.encode(y["query_response0"]))259 y["query_response1"] = y["query"].strip() + y["response1"]260 y["query_response1_token"] = tokenizer.encode(261 y["query_response1"], padding="max_length", max_length=args.max_rm_query_response_length, truncation=True262 )263 y["query_response1_token_len"] = len(tokenizer.encode(y["query_response1"]))264 return y265 266 label_ds = label_ds.map(process_response_data, load_from_cache_file=False, num_proc=multiprocessing.cpu_count())267 if args.push_to_hub:268 label_ds.push_to_hub(f"{args.hf_entity}/summarize_from_feedback_oai_preprocessing_{timestamp}")269 270 os.makedirs("dataset_visuals", exist_ok=True)271 # visualize token length distribution272 num_subplots = len(sft_ds) * 2 + len(label_ds) * 4273 print(f"{num_subplots=}")274 fig, axs = plt.subplots(5, 3, figsize=(16, 16))275 axs = axs.flatten()276 j = 0277 for _, key in enumerate(sft_ds.keys()):278 df = sft_ds[key].to_pandas()279 axs[j].hist(df["reference_response_token_len"], bins=100)280 axs[j].set_title(f"{key} split: reference response token length\nmax_length={max(df['reference_response_token_len'])}")281 axs[j + 1].hist(df["query_reference_response_token_len"], bins=100)282 axs[j + 1].set_title(283 f"{key} split: query.strip() + reference response token length\nmax_length={max(df['query_reference_response_token_len'])}"284 )285 j += 2286 offset = len(sft_ds)287 for _, key in enumerate(label_ds.keys()):288 df = label_ds[key].to_pandas()289 axs[j].hist(df["response0_token_len"], bins=100)290 axs[j].set_title(f"{key} split: response0 token length\nmax_length={max(df['response0_token_len'])}")291 axs[j + 1].hist(df["response1_token_len"], bins=100)292 axs[j + 1].set_title(f"{key} split: response1 token length\nmax_length={max(df['response1_token_len'])}")293 axs[j + 2].hist(df["query_response0_token_len"], bins=100)294 axs[j + 2].set_title(295 f"{key} split: query.strip() + response0 token length\nmax_length={max(df['query_response0_token_len'])}"296 )297 axs[j + 3].hist(df["query_response1_token_len"], bins=100)298 axs[j + 3].set_title(299 f"{key} split: query.strip() + response1 token length\nmax_length={max(df['query_response1_token_len'])}"300 )301 j += 4302 fig.suptitle(f"{args.base_model} Tokenizer: Token length distribution")303 fig.tight_layout()304 fig.savefig("dataset_visuals/token_len.png")305 306 # visualize confidence distribution307 fig, axs = plt.subplots(len(label_ds), 1, figsize=(8, 8))308 axs = axs.flatten()309 label_ds = label_ds.flatten()310 for i, key in enumerate(label_ds.keys()):311 df = label_ds[key].to_pandas()312 axs[i].hist(df["extra.confidence"])313 axs[i].set_title(f"{key} split: confidence distribution")314 fig.suptitle("Confidence distribution")315 fig.tight_layout()316 fig.savefig("dataset_visuals/confidence.png")317 318 # visualize policies used319 fig, axs = plt.subplots(1, len(label_ds), figsize=(8, 12))320 axs = axs.flatten()321 label_ds = label_ds.flatten()322 for i, key in enumerate(label_ds.keys()):323 df = label_ds[key].to_pandas()324 cat = pd.concat([df["response0_policy"], df["response1_policy"]], axis=0)325 cat.hist(ax=axs[i], xrot=90, orientation="horizontal")326 axs[i].set_title(f"{key} split: policy distribution")327 fig.suptitle("Policy distribution")328 fig.tight_layout()329 fig.savefig("dataset_visuals/policies.png")330 331 # visualize compairson distribution332 fig, axs = plt.subplots(1, len(label_ds), figsize=(24, 30))333 axs = axs.flatten()334 label_ds = label_ds.flatten()335 for i, key in enumerate(label_ds.keys()):336 df = label_ds[key].to_pandas()337 df["policies"].hist(ax=axs[i], xrot=90, orientation="horizontal")338 axs[i].set_title(f"{key} split: policy comparison distribution")339 fig.suptitle("Policy comparison distribution")340 fig.tight_layout()341 fig.savefig("dataset_visuals/policy_comparisons.png")342 343 if args.push_to_hub:344 # upload the `dataset_visuals`345 api.upload_folder(346 folder_path="dataset_visuals",347 path_in_repo="dataset_visuals",348 repo_id=f"{args.hf_entity}/summarize_from_feedback_oai_preprocessing_{timestamp}",349 repo_type="dataset",350 )351 # upload current file352 print(f"{__file__=}")353 api.upload_file(354 path_or_fileobj=__file__,355 path_in_repo="create_dataset.py",356 repo_id=f"{args.hf_entity}/summarize_from_feedback_oai_preprocessing_{timestamp}",357 repo_type="dataset",358 )359 