Views
No views yet
Apache license 2.0| Metrics | Base | LoRA |
|---|---|---|
| BLEU | 16.10 | 38.95 |
| ROUGE-1 | 0.299 | 0.323 |
| ROUGE-2 | 0.090 | 0.091 |
| ROUGE-L | 0.218 | 0.230 |
!pip install -U bitsandbytes transformers peft accelerate1import re, torch
2from typing import List, Optional
3from transformers import AutoTokenizer, AutoModelForCausalLM
4from transformers import BitsAndBytesConfig
5from peft import PeftModel, PeftConfig
6
7_THINK_RE = re.compile(r"(?is)<think>.*?</think>")
8_LEAD_ASSIST = re.compile(r"(?is)^.*?\bassistant\b[:\-]?\s*")
9
10def clean_text(s: str) -> str:
11 s = _THINK_RE.sub("", s)
12 s = _LEAD_ASSIST.sub("", s.strip(), count=1)
13 return s.strip()
14
15def load_model_and_tokenizer(
16 base_model: Optional[str],
17 adapter_repo: str,
18 max_seq_len: int = 4096,
19 load_in_4bit: bool = True,
20 compute_dtype: torch.dtype = torch.bfloat16,
21):
22 peft_cfg = PeftConfig.from_pretrained(adapter_repo)
23 suggested_base = getattr(peft_cfg, "base_model_name_or_path", None)
24
25 if base_model is None:
26 base_model = suggested_base
27 print(f"[info] Using base from adapter config: {base_model}")
28 elif suggested_base and (base_model != suggested_base):
29 print(f"[warn] Adapter expects base '{suggested_base}', "
30 f"but you set '{base_model}'. Make sure they match!")
31
32 quant_cfg = None
33 if load_in_4bit:
34 quant_cfg = BitsAndBytesConfig(
35 load_in_4bit=True,
36 bnb_4bit_use_double_quant=True,
37 bnb_4bit_compute_dtype=compute_dtype,
38 bnb_4bit_quant_type="nf4",
39 )
40
41 tok = AutoTokenizer.from_pretrained(base_model, use_fast=True, trust_remote_code=True)
42
43 if tok.pad_token_id is None and tok.eos_token_id is not None:
44 tok.pad_token = tok.eos_token
45 tok.pad_token_id = tok.eos_token_id
46
47 tok.padding_side = "left"
48 tok.truncation_side = "left"
49
50 model = AutoModelForCausalLM.from_pretrained(
51 base_model,
52 device_map="auto",
53 torch_dtype="auto" if not load_in_4bit else None,
54 quantization_config=quant_cfg,
55 trust_remote_code=True,
56 )
57
58 model = PeftModel.from_pretrained(model, adapter_repo)
59
60 if getattr(model, "generation_config", None) is not None and tok.pad_token_id is not None:
61 model.generation_config.pad_token_id = tok.pad_token_id
62
63 return model, tok
64
65def build_chat(tok, user_text: str, system_text: Optional[str] = None) -> str:
66 messages = []
67 if system_text:
68 messages.append({"role": "system", "content": system_text})
69 messages.append({"role": "user", "content": user_text})
70 prompt = tok.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
71 return prompt
72
73@torch.inference_mode()
74def generate_one(
75 model,
76 tok,
77 user_text: str,
78 system_text: str = "Answer concisely, straight to the point, no <think>.",
79 max_new_tokens: int = 200,
80 temperature: float = 0.7,
81 top_p: float = 0.9,
82):
83 prompt = build_chat(tok, user_text, system_text)
84 device = next(model.parameters()).device
85 inputs = tok(prompt, return_tensors="pt").to(device)
86 in_len = inputs["input_ids"].shape[1]
87
88 im_end_id = tok.convert_tokens_to_ids("<|im_end|>")
89 eos_ids = [i for i in [tok.eos_token_id, im_end_id] if i is not None]
90 eos_ids = eos_ids[0] if len(eos_ids) == 1 else eos_ids
91
92 out = model.generate(
93 **inputs,
94 max_new_tokens=max_new_tokens,
95 do_sample=(temperature is not None and temperature > 0),
96 temperature=temperature,
97 top_p=top_p,
98 num_beams=1,
99 eos_token_id=eos_ids,
100 pad_token_id=tok.pad_token_id,
101 use_cache=True,
102 no_repeat_ngram_size=3,
103 repetition_penalty=1.1,
104 )[0]
105
106 gen_ids = out[in_len:]
107 text = tok.decode(gen_ids, skip_special_tokens=True)
108 return clean_text(text)
109
110@torch.inference_mode()
111def generate_batch(
112 model,
113 tok,
114 user_texts: List[str],
115 system_text: Optional[str] = None,
116 max_new_tokens: int = 200,
117 temperature: float = 0.7,
118 top_p: float = 0.9,
119 batch_size: int = 8,
120):
121 device = next(model.parameters()).device
122 im_end_id = tok.convert_tokens_to_ids("<|im_end|>")
123 eos_ids = [i for i in [tok.eos_token_id, im_end_id] if i is not None]
124 eos_ids = eos_ids[0] if len(eos_ids) == 1 else eos_ids
125
126 answers = []
127 for i in range(0, len(user_texts), batch_size):
128 chunk = user_texts[i:i + batch_size]
129 prompts = [build_chat(tok, u, system_text) for u in chunk]
130 toks = tok(prompts, return_tensors="pt", padding=True).to(device)
131 in_lens = toks["attention_mask"].sum(dim=1).tolist()
132
133 outs = model.generate(
134 **toks,
135 max_new_tokens=max_new_tokens,
136 do_sample=(temperature is not None and temperature > 0),
137 temperature=temperature,
138 top_p=top_p,
139 num_beams=1,
140 eos_token_id=eos_ids,
141 pad_token_id=tok.pad_token_id,
142 use_cache=True,
143 no_repeat_ngram_size=3,
144 repetition_penalty=1.1,
145 )
146
147 for out, L in zip(outs, in_lens):
148 ans = tok.decode(out[L:], skip_special_tokens=True)
149 answers.append(clean_text(ans))
150 return answers
151
152if __name__ == "__main__":
153 adapter = "luminolous/astropher-lora"
154 base = "unsloth/Qwen3-1.7B-unsloth-bnb-4bit"
155
156 model, tok = load_model_and_tokenizer(
157 base_model=base,
158 adapter_repo=adapter,
159 max_seq_len=4096,
160 load_in_4bit=True,
161 compute_dtype=torch.bfloat16,
162 )
163
164 q = "What is inside a black hole?" # <- You can change the question here
165 print(f"\nModel output: {generate_one(model, tok, q, max_new_tokens=180)}")
166