CoolFace
Datasetpublic

cleanrl/summarize_from_feedback_oai_preprocessing_1704321749

Dataset Card for "summarize_from_feedback_oai_preprocessing_1704321749" More Information needed

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes69downloads
create_dataset.py359 linesDownload Raw Back to root
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