Team Ai
Apppublic

alirezaaminzadeh/python-docstring-generator

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
inference.py30 linesDownload Raw Back to root
1"""
2Inference for docstring generation. Uses T5 (cached after first load).
3"""
4import torch
5
6_cache = {}
7
8def generate_docstring(
9    code: str,
10    model_name: str = "t5-small",
11    max_length: int = 128,
12    num_beams: int = 4,
13    device: str = None,
14) -> str:
15    if device is None:
16        device = "cuda" if torch.cuda.is_available() else "cpu"
17    if model_name not in _cache:
18        from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
19        _cache[model_name] = {
20            "tokenizer": AutoTokenizer.from_pretrained(model_name),
21            "model": AutoModelForSeq2SeqLM.from_pretrained(model_name).to(device),
22        }
23    tokenizer = _cache[model_name]["tokenizer"]
24    model = _cache[model_name]["model"]
25    input_text = "summarize: " + code
26    inputs = tokenizer(input_text, return_tensors="pt", truncation=True, max_length=512).to(device)
27    with torch.no_grad():
28        out = model.generate(**inputs, max_length=max_length, num_beams=num_beams, early_stopping=True)
29    return tokenizer.decode(out[0], skip_special_tokens=True)
30