Team Ai
Apppublic

tmnam20/code-summarization

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
st_utils.py237 linesDownload Raw Back to root
1from __future__ import absolute_import2import streamlit as st3import torch4import os5import sys6import pickle7import torch8import json9import random10import logging11import argparse12import numpy as np13from io import open14from itertools import cycle15import torch.nn as nn16from model import Seq2Seq17from tqdm import tqdm, trange18import regex as re19from torch.utils.data import (20    DataLoader,21    Dataset,22    SequentialSampler,23    RandomSampler,24    TensorDataset,25)26from torch.utils.data.distributed import DistributedSampler27from transformers import (28    WEIGHTS_NAME,29    AdamW,30    get_linear_schedule_with_warmup,31    RobertaConfig,32    RobertaModel,33    RobertaTokenizer,34)35from huggingface_hub import hf_hub_download36import io37 38# def list_files(startpath, prev_level=0):39#     # list files recursively40#     for root, dirs, files in os.walk(startpath):41#         level = root.replace(startpath, "").count(os.sep) + prev_level42#         indent = " " * 4 * (level)43 44#         print("{}{}/".format(indent, os.path.basename(root)))45#         # st.write("{}{}/".format(indent, os.path.basename(root)))46 47#         subindent = " " * 4 * (level + 1)48#         for f in files:49#             print("{}{}".format(subindent, f))50#             # st.write("{}{}".format(subindent, f))51 52#         for d in dirs:53#             list_files(d, level + 1)54 55 56class CONFIG:57    max_source_length = 25658    max_target_length = 12859    beam_size = 360    local_rank = -161    no_cuda = False62 63    do_train = True64    do_eval = True65    do_test = True66    train_batch_size = 1267    eval_batch_size = 3268 69    model_type = "roberta"70    model_name_or_path = "microsoft/codebert-base"71    output_dir = "/content/drive/MyDrive/CodeSummarization"72    load_model_path = None73    train_filename = "dataset/python/train.jsonl"74    dev_filename = "dataset/python/valid.jsonl"75    test_filename = "dataset/python/test.jsonl"76    config_name = ""77    tokenizer_name = ""78    cache_dir = "cache"79 80    save_every = 500081 82    gradient_accumulation_steps = 183    learning_rate = 5e-584    weight_decay = 1e-485    adam_epsilon = 1e-886    max_grad_norm = 1.087    num_train_epochs = 3.088    max_steps = -189    warmup_steps = 090    train_steps = 10000091    eval_steps = 1000092    n_gpu = torch.cuda.device_count()93 94 95# download model with streamlit cache decorator96@st.cache_resource97def download_model():98    if not os.path.exists(r"models/pytorch_model.bin"):99        os.makedirs("./models", exist_ok=True)100        path = hf_hub_download(101            repo_id="tmnam20/codebert-code-summarization",102            filename="pytorch_model.bin",103            cache_dir="cache",104            local_dir=os.path.join(os.getcwd(), "models"),105            local_dir_use_symlinks=False,106            force_download=True,107        )108 109 110# load with streamlit cache decorator111# @st.cache(persist=False, show_spinner=True, allow_output_mutation=True)112@st.cache_resource113def load_tokenizer_and_model(pretrained_path):114    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")115 116    # Config model117    config_class, model_class, tokenizer_class = (118        RobertaConfig,119        RobertaModel,120        RobertaTokenizer,121    )122    model_config = config_class.from_pretrained(123        CONFIG.config_name if CONFIG.config_name else CONFIG.model_name_or_path,124        cache_dir=CONFIG.cache_dir,125    )126    # model_config.save_pretrained("config")127 128    # load tokenizer129    tokenizer = tokenizer_class.from_pretrained(130        CONFIG.tokenizer_name if CONFIG.tokenizer_name else CONFIG.model_name_or_path,131        cache_dir=CONFIG.cache_dir,132        # do_lower_case=args.do_lower_case133    )134 135    # load encoder from pretrained RoBERTa136    encoder = model_class.from_pretrained(137        CONFIG.model_name_or_path, config=model_config, cache_dir=CONFIG.cache_dir138    )139 140    # build decoder141    decoder_layer = nn.TransformerDecoderLayer(142        d_model=model_config.hidden_size, nhead=model_config.num_attention_heads143    )144    decoder = nn.TransformerDecoder(decoder_layer, num_layers=6)145 146    # build seq2seq model from pretrained encoder and from-scratch decoder147    model = Seq2Seq(148        encoder=encoder,149        decoder=decoder,150        config=model_config,151        beam_size=CONFIG.beam_size,152        max_length=CONFIG.max_target_length,153        sos_id=tokenizer.cls_token_id,154        eos_id=tokenizer.sep_token_id,155    )156 157    try:158        state_dict = torch.load(159            os.path.join(os.getcwd(), "models", "pytorch_model.bin"),160            map_location=device,161        )162    except RuntimeError as e:163        print(e)164        try:165            state_dict = torch.load(166                os.path.join(os.getcwd(), "models", "pytorch_model.bin"),167                map_location="cpu",168            )169        except RuntimeError as e:170            print(e)171            state_dict = torch.load(172                os.path.join(os.getcwd(), "models", "pytorch_model_cpu.bin"),173                map_location="cpu",174            )175 176    del state_dict["encoder.embeddings.position_ids"]177    model.load_state_dict(state_dict)178 179    # model = model.to("cpu")180    # torch.save(model.state_dict(), os.path.join(os.getcwd(), "models", "pytorch_model_cpu.bin"))181 182    model = model.to(device)183 184    return tokenizer, model, device185 186 187@st.cache_data188def preprocessing(code_segment):189    # remove newlines190    code_segment = re.sub(r"\n", " ", code_segment)191 192    # remove docstring193    code_segment = re.sub(r'""".*?"""', "", code_segment, flags=re.DOTALL)194 195    # remove multiple spaces196    code_segment = re.sub(r"\s+", " ", code_segment)197 198    # remove comments199    code_segment = re.sub(r"#.*", "", code_segment)200 201    # remove html tags202    code_segment = re.sub(r"<.*?>", "", code_segment)203 204    # remove urls205    code_segment = re.sub(r"http\S+", "", code_segment)206 207    # split special chars into different tokens208    code_segment = re.sub(r"([^\w\s])", r" \1 ", code_segment)209 210    return code_segment.split()211 212 213def generate_docstring(model, tokenizer, device, code_segemnt, max_length=None):214    input_tokens = preprocessing(code_segemnt)215    encoded_input = tokenizer.encode_plus(216        input_tokens,217        max_length=CONFIG.max_source_length,218        pad_to_max_length=True,219        truncation=True,220        return_tensors="pt",221    )222 223    input_ids = encoded_input["input_ids"].to(device)224    input_mask = encoded_input["attention_mask"].to(device)225 226    if max_length is not None:227        model.max_length = max_length228 229    summary = model(input_ids, input_mask)230 231    # decode summary with tokenizer232    summaries = []233    for i in range(summary.shape[1]):234        summaries.append(tokenizer.decode(summary[0][i], skip_special_tokens=True))235    return summaries236    # return tokenizer.decode(summary[0][0], skip_special_tokens=True)237