62 from peft
import LoraConfig, TaskType, get_peft_model
63 from transformers
import (
66 DataCollatorForLanguageModeling,
70 except ImportError
as exc:
72 "LLM LoRA train needs transformers, peft, datasets, and torch."
75 data_path = Path(str(recipe.data[
"path"]))
77 text_field = str(recipe.data.get(
"text_field")
or "text")
78 texts = [
_row_text(r, text_field)
for r
in rows]
80 base = str(recipe.train[
"base"])
81 rank = int(recipe.train.get(
"rank")
or 8)
82 lora_alpha = int(recipe.train.get(
"lora_alpha")
or 16)
83 lr = float(recipe.train.get(
"lr")
or 2e-4)
84 epochs = float(recipe.train.get(
"epochs")
or 1)
85 max_steps = recipe.train.get(
"max_steps")
86 max_seq = int(recipe.train.get(
"max_seq_len")
or 512)
87 batch = int(recipe.train.get(
"batch_size")
or 1)
88 accum = int(recipe.train.get(
"grad_accum")
or 4)
90 tokenizer = AutoTokenizer.from_pretrained(base, use_fast=
True)
91 if tokenizer.pad_token
is None:
92 tokenizer.pad_token = tokenizer.eos_token
94 model = AutoModelForCausalLM.from_pretrained(base)
95 model.config.pad_token_id = tokenizer.pad_token_id
96 lora_kwargs: dict[str, Any] = dict(
97 task_type=TaskType.CAUSAL_LM,
99 lora_alpha=lora_alpha,
100 lora_dropout=float(recipe.train.get(
"dropout")
or 0.05),
102 modules = recipe.train.get(
"target_modules")
104 lora_kwargs[
"target_modules"] = modules
105 model = get_peft_model(model, LoraConfig(**lora_kwargs))
107 ds = Dataset.from_dict({
"text": texts})
109 def tokenize(batch: dict[str, list[str]]) -> dict[str, Any]:
117 tokenized = ds.map(tokenize, batched=
True, remove_columns=[
"text"])
118 ckpt_dir = out_dir /
"checkpoints" /
"lora"
119 ckpt_dir.mkdir(parents=
True, exist_ok=
True)
121 args = TrainingArguments(
122 output_dir=str(out_dir /
"hf_trainer"),
123 per_device_train_batch_size=batch,
124 gradient_accumulation_steps=accum,
126 num_train_epochs=epochs,
127 max_steps=int(max_steps)
if max_steps
is not None else -1,
131 fp16=bool(torch.cuda.is_available()),
133 remove_unused_columns=
False,
135 collator = DataCollatorForLanguageModeling(tokenizer, mlm=
False)
136 losses: list[float] = []
139 def log(self, logs: dict[str, float], *rest: Any, **kwargs: Any) ->
None:
140 super().log(logs, *rest, **kwargs)
142 losses.append(float(logs[
"loss"]))
143 recorder.log(int(self.state.global_step
or 0), loss=float(logs[
"loss"]))
148 train_dataset=tokenized,
149 data_collator=collator,
152 model.save_pretrained(str(ckpt_dir))
153 tokenizer.save_pretrained(str(ckpt_dir))
154 recorder.signal(int(trainer.state.global_step
or 1),
"lora checkpoint saved")
156 mean_loss = float(sum(losses) / len(losses))
if losses
else None
157 metrics: dict[str, Any] = {
"train_loss": mean_loss,
"n_rows": len(texts)}
158 inspect = {
"n_rows": len(texts),
"base": base,
"rank": rank}
160 probes_path = recipe.eval.get(
"probes")
161 min_score = recipe.eval.get(
"min_score")
162 details: dict[str, Any] = {
"train_loss": mean_loss}
166 scores: list[float] = []
167 rows_out: list[dict[str, Any]] = []
168 max_new = int(recipe.eval.get(
"max_tokens")
or 32)
170 prompt = str(p.get(
"prompt")
or p.get(
"text")
or "")
171 ref = str(p.get(
"reference")
or p.get(
"completion")
or "")
172 inputs = tokenizer(prompt, return_tensors=
"pt")
173 inputs = {k: v.to(model.device)
for k, v
in inputs.items()}
174 with torch.no_grad():
175 out = model.generate(**inputs, max_new_tokens=max_new, do_sample=
False)
176 pred = tokenizer.decode(out[0][inputs[
"input_ids"].shape[1] :], skip_special_tokens=
True)
177 sc =
_overlap(pred, ref)
if ref
else 0.0
179 rows_out.append({
"prompt": prompt,
"response": pred,
"score": sc})
180 score = float(sum(scores) / len(scores))
if scores
else 0.0
181 metric = str(recipe.eval.get(
"metric")
or "overlap")
182 skipped = min_score
is None
183 passed =
True if skipped
else score >= float(min_score)
184 details[
"probes"] = rows_out
188 score=round(score, 6),
189 min_score=
None if min_score
is None else float(min_score),
193 metrics[
"eval_overlap"] = score
195 metric =
"train_loss"
196 skipped = min_score
is None
197 if skipped
or mean_loss
is None:
203 passed = mean_loss <= float(min_score)
207 score=
None if score
is None else round(float(score), 6),
208 min_score=
None if min_score
is None else float(min_score),
217 step=int(trainer.state.global_step
or 1),
218 extra={
"base": base,
"rank": rank},