Views
No views yet
| Metric | Base Gemma-2B | SFT | Δ |
|---|---|---|---|
| Perplexity | 46.88 | 17.11 | ↓ 29.77 |
| ROUGE-1 | 14.95 | 25.16 | +10.21 |
| ROUGE-2 | 3.81 | 9.30 | +5.49 |
| ROUGE-L | 12.78 | 21.52 | +8.74 |
| BERTScore-F1 | 69.63 | 75.40 | +5.77 |
1from transformers import AutoTokenizer, AutoModelForCausalLM, BitsAndBytesConfig
2from peft import PeftModel
3import torch
4
5MODEL_ID = "google/gemma-2b"
6ADAPTER = "Yash983/medgemma-2b-sft"
7
8# ── 4-bit quantization — without this, 2B params load as float32 (~16GB RAM) ──
9bnb = BitsAndBytesConfig(
10 load_in_4bit=True,
11 bnb_4bit_use_double_quant=True,
12 bnb_4bit_quant_type="nf4",
13 bnb_4bit_compute_dtype=torch.float16,
14)
15
16tokenizer = AutoTokenizer.from_pretrained(MODEL_ID)
17tokenizer.pad_token = tokenizer.eos_token
18
19model = AutoModelForCausalLM.from_pretrained(
20 MODEL_ID,
21 quantization_config=bnb,
22 device_map="auto",
23 torch_dtype=torch.float16,
24)
25model = PeftModel.from_pretrained(model, ADAPTER)
26model.eval() # disable dropout
27
28prompt = "Explain symptoms of diabetes"
29inputs = tokenizer(prompt, return_tensors="pt").to(model.device) # must match model device
30
31with torch.no_grad(): # no gradient tracking at inference
32 outputs = model.generate(
33 **inputs,
34 max_new_tokens=200,
35 do_sample=True, # required when using temperature
36 temperature=0.7,
37 top_p=0.9, # pair with temperature for quality
38 pad_token_id=tokenizer.eos_token_id,
39 eos_token_id=tokenizer.eos_token_id,
40 )
41
42# strip prompt tokens — otherwise the prompt is printed again at the start
43generated_ids = outputs[0][inputs["input_ids"].shape[1]:]
44print(tokenizer.decode(generated_ids, skip_special_tokens=True))1@misc{medgemma2b-sft,
2 title = {MedGemma-2B-SFT: Supervised fine-tuned medical language model},
3 author = {Yash},
4 year = {2025},
5 base = {google/gemma-2b},
6 url = {https://huggingface.co/Yash983/medgemma-2b-sft}
7}