{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"%%writefile prepare_hf_dataset.py\n\nfrom functools import partial\nfrom tqdm.auto import tqdm\nimport pandas as pd\nimport numpy as np\nimport argparse\nimport gc\n\nimport transformers\nimport datasets\n\ntqdm.pandas()\nCELL_SEP = '[CELL_SEP]'\n\ndef prune_code_tokens(code_token_ids, max_seq_len):\n    \"\"\"\n    Prunes cells that take too many tokens to fit in max_seq_len.\n    \"\"\"\n    code_token_counts = [len(token_ids) for token_ids in code_token_ids]\n    total_number_of_cells = len(code_token_counts)\n    total_tokens_to_prune = max(sum(code_token_counts)-max_seq_len, 0)\n\n    tokens_to_prune_per_cell = [0]*total_number_of_cells\n    total_pruned_tokens = 0\n    while total_tokens_to_prune > 0:\n        cur_max_code_token_count = max(code_token_counts)\n        second_max_code_token_count = sorted(code_token_counts)[-2]\n        for cell_idx, code_token_count in enumerate(code_token_counts):\n            if not code_token_count == cur_max_code_token_count: \n                continue\n            \n            num_tokens_to_pop = min(code_token_count-second_max_code_token_count+1, total_tokens_to_prune)\n            tokens_to_prune_per_cell[cell_idx] += num_tokens_to_pop\n            total_pruned_tokens += num_tokens_to_pop\n            total_tokens_to_prune -= num_tokens_to_pop\n            code_token_counts[cell_idx] -= num_tokens_to_pop\n            break\n    \n    # Prune the cell tokens\n    pruned_code_token_ids = []\n    for code_token_ids, num_tokens_to_pop in zip(code_token_ids, tokens_to_prune_per_cell):\n        if num_tokens_to_pop == 0:\n            pruned_code_token_ids.append(code_token_ids)\n            continue\n        pruned_code_token_ids.append(code_token_ids[:-num_tokens_to_pop])\n    return pruned_code_token_ids\n\ndef shift_tokens_right(input_ids, pad_token_id, decoder_start_token_id):\n    \"\"\" Shift input ids one token to the right \"\"\"\n    shifted_input_ids = np.zeros_like(input_ids)\n    shifted_input_ids[1:] = input_ids[:-1]\n    shifted_input_ids[0] = decoder_start_token_id\n    shifted_input_ids = np.where(shifted_input_ids == -100, pad_token_id, shifted_input_ids)\n    return shifted_input_ids\n\ndef convert_to_features_seq2seq(\n    notebook_dict,\n    tokenizer,\n    max_input_seq_len,\n    max_target_seq_len,\n    max_markdown_seq_len,\n    max_tokens_per_cell,\n):\n    '''Tokenize the notebook and convert to features for the model'''\n\n    markdown_cell_sources = notebook_dict['merged_markdown_cell_sources'].split(CELL_SEP)\n    markdown_cell_pct_ranks = [float(rank) for rank in notebook_dict['merged_markdown_cell_pct_ranks'].split(CELL_SEP)]\n    markdown_cell_ids = notebook_dict['merged_markdown_cell_ids'].split(CELL_SEP)\n    markdown_special_tokens = [f'M{cell_idx} ' for cell_idx in range(len(markdown_cell_ids))]\n    markdown_cell_sources = [sp_token+source for sp_token, source in zip(markdown_special_tokens, markdown_cell_sources)]\n\n    code_cell_sources = notebook_dict['merged_code_cell_sources'].split(CELL_SEP)\n    code_cell_pct_ranks = [float(rank) for rank in notebook_dict['merged_code_cell_pct_ranks'].split(CELL_SEP)]\n    code_cell_ids = notebook_dict['merged_code_cell_ids'].split(CELL_SEP)\n    code_special_tokens = [f'C{cell_idx} ' for cell_idx in range(len(code_cell_ids))]\n    code_cell_sources = [sp_token+source for sp_token, source in zip(code_special_tokens, code_cell_sources)]\n\n    # Remove cells from the end of the notebook so that all cells have at least one representative token\n    max_markdown_cells = max_markdown_seq_len//2\n    max_code_cells = (max_input_seq_len-max_markdown_seq_len)//2\n    if len(markdown_cell_sources) > max_markdown_cells:\n        markdown_cell_sources = markdown_cell_sources[:max_markdown_cells]\n        markdown_cell_pct_ranks = markdown_cell_pct_ranks[:max_markdown_cells]\n        markdown_cell_ids = markdown_cell_ids[:max_markdown_cells]\n        markdown_special_tokens = markdown_special_tokens[:max_markdown_cells]\n    if len(code_cell_sources) > max_code_cells:\n        code_cell_sources = code_cell_sources[:max_code_cells]\n        code_cell_pct_ranks = code_cell_pct_ranks[:max_code_cells]\n        code_cell_ids = code_cell_ids[:max_code_cells]\n        code_special_tokens = code_special_tokens[:max_code_cells]\n    \n    markdown_cell_count = len(markdown_cell_sources)\n    code_cell_count = len(code_cell_sources)\n\n    max_tokens_per_markdown_cell = max(max_tokens_per_cell, max_markdown_seq_len//markdown_cell_count)\n    markdown_cell_token_ids = tokenizer(\n        markdown_cell_sources,\n        max_length=max_tokens_per_markdown_cell,\n        truncation=True,\n    )['input_ids']\n    markdown_cell_token_ids = prune_code_tokens(markdown_cell_token_ids, max_markdown_seq_len)\n    total_markdown_code_tokens = sum([len(token_ids) for token_ids in markdown_cell_token_ids])\n\n    max_code_seq_len = max_input_seq_len - total_markdown_code_tokens\n    max_tokens_per_code_cell = max(max_tokens_per_cell, max_code_seq_len//code_cell_count)\n    code_cell_token_ids = tokenizer(\n        code_cell_sources, \n        max_length=max_tokens_per_code_cell, \n        truncation=True, \n    )['input_ids']\n    code_cell_token_ids = prune_code_tokens(code_cell_token_ids, max_input_seq_len-total_markdown_code_tokens)\n\n\n    # Map cell_pct_rank -> special token\n    all_special_tokens = markdown_special_tokens + code_special_tokens\n    all_cell_pct_ranks = markdown_cell_pct_ranks + code_cell_pct_ranks\n    cell_pct_rank_to_special_token = {\n        cell_pct_rank: special_token \n        for special_token, cell_pct_rank in zip(all_special_tokens, all_cell_pct_ranks)\n    }\n\n    # Get model output according to sorted global rank for special tokens\n    sorted_cell_pct_ranks = sorted(markdown_cell_pct_ranks + code_cell_pct_ranks)\n    model_output_special_tokens = [\n        cell_pct_rank_to_special_token[cell_pct_rank]\n        for cell_pct_rank in sorted_cell_pct_ranks\n    ]\n    \n    with tokenizer.as_target_tokenizer():\n        model_output_str = ''.join(model_output_special_tokens)\n        target_ids = tokenizer(\n            model_output_str,\n            max_length=max_target_seq_len,\n            padding='max_length',\n            truncation=True,\n        )['input_ids']\n        target_ids = shift_tokens_right(target_ids, tokenizer.pad_token_id, tokenizer.pad_token_id)\n    \n    input_ids = []\n    for cell_ids in markdown_cell_token_ids + code_cell_token_ids:\n        input_ids += cell_ids\n    num_input_pad_tokens = max_input_seq_len-len(input_ids)\n    attention_mask = [1]*len(input_ids) + [0]*num_input_pad_tokens\n    input_ids += [0]*num_input_pad_tokens\n    \n    notebook_features = {\n        'input_ids': input_ids, \n        'attention_mask': attention_mask,\n        'target_ids': target_ids,\n        'notebook_id': notebook_dict['notebook_id'],\n    }\n    # print(notebook_features)\n    return notebook_features\n\n\ndef build_hf_dataset(\n    df, \n    tokenizer, \n    max_input_seq_len,\n    max_target_seq_len,\n    max_markdown_seq_len,\n    max_tokens_per_cell,\n    ):\n    '''Builds the huggingface dataset for training the model.'''\n    convert_to_features = partial(\n        convert_to_features_seq2seq, \n        tokenizer=tokenizer,\n        max_input_seq_len=max_input_seq_len,\n        max_target_seq_len=max_target_seq_len,\n        max_markdown_seq_len=max_markdown_seq_len,\n        max_tokens_per_cell=max_tokens_per_cell,\n    )\n    raw_dataset = datasets.Dataset.from_pandas(df)\n    processed_dataset = raw_dataset.map(\n        convert_to_features, \n        remove_columns=raw_dataset.column_names, \n        desc='Running tokenizer on raw dataset'\n    )\n    processed_dataset.set_format(type='numpy')\n    empty_sentences = (np.array(processed_dataset['attention_mask'])[:, -1] == 0).sum()\n    print('Empty sentences ratio:', empty_sentences/len(processed_dataset))\n    return processed_dataset\n\nif __name__ == '__main__':\n    parser = argparse.ArgumentParser()\n    parser.add_argument('--tokenizer_name', default='google/bigbird-roberta-large', type=str, help='The tokenizer name')\n    parser.add_argument('--max_input_seq_length', default=1024, type=int, help='The max sequence length for the input')\n    parser.add_argument('--max_target_seq_length', default=1024, type=int, help='The max sequence length for the target')\n    parser.add_argument('--max_markdown_seq_length', default=512, type=int, help='The max markdown sequence length')\n    parser.add_argument('--max_tokens_per_cell', default=128, type=int, help='The max tokens per cell')\n    parser.add_argument('--notebooks_df_path', default='notebooks_df.csv', type=str, help='Path to notebooks.csv')\n\n    args = parser.parse_args()\n    \n    tokenizer = transformers.AutoTokenizer.from_pretrained(args.tokenizer_name)\n    notebooks_df = pd.read_csv(args.notebooks_df_path)\n    print('Total number of notebooks:', len(notebooks_df))\n    \n    for fold in tqdm(range(8), desc='Tokenizing notebooks for each fold'):\n        fold_df = notebooks_df[notebooks_df.notebook_fold == fold]\n        fold_dataset = build_hf_dataset(\n            df=fold_df,\n            tokenizer=tokenizer,\n            max_input_seq_len=args.max_input_seq_length,\n            max_target_seq_len=args.max_target_seq_length,\n            max_markdown_seq_len=args.max_markdown_seq_length,\n            max_tokens_per_cell=args.max_tokens_per_cell,\n        )\n        fold_dataset.save_to_disk(f'hf_dataset_fold_{fold}')\n    print('Done!')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q --upgrade transformers\n!python prepare_hf_dataset.py \\\n    --tokenizer_name 'google/long-t5-tglobal-base' \\\n    --max_input_seq_len 2048 \\\n    --max_target_seq_len 512 \\\n    --max_markdown_seq_len 1024 \\\n    --max_tokens_per_cell 256 \\\n    --notebooks_df_path '/kaggle/input/ai4code-dataframes/notebooks_df.csv'","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}