alirezaaminzadeh/python-docstring-generator
0
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 