-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathtrain.py
More file actions
121 lines (96 loc) · 3.59 KB
/
Copy pathtrain.py
File metadata and controls
121 lines (96 loc) · 3.59 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
"""Finetune script template for functiongemma / qwen3 (real training path)."""
from __future__ import annotations
import argparse
import json
import os
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Dict, List, Optional
from datasets import load_dataset
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
Trainer,
TrainingArguments,
DataCollatorForLanguageModeling,
)
try:
from peft import LoraConfig, get_peft_model, TaskType
except Exception as exc: # pragma: no cover
raise RuntimeError("peft is required for LoRA finetuning") from exc
@dataclass
class TrainConfig:
model_id: str
dataset_path: str
output_dir: str
num_train_epochs: int = 2
per_device_train_batch_size: int = 2
gradient_accumulation_steps: int = 4
learning_rate: float = 2e-5
max_seq_length: int = 512
lora_r: int = 16
lora_alpha: int = 32
lora_dropout: float = 0.05
seed: int = 42
text_field: Optional[str] = None
def load_config(path: str) -> TrainConfig:
raw = json.loads(Path(path).read_text(encoding="utf-8"))
return TrainConfig(**raw)
def format_sample(sample: Dict[str, Any]) -> str:
if "text" in sample and sample["text"]:
return str(sample["text"])
instruction = str(sample.get("instruction") or "")
input_text = str(sample.get("input") or "")
output_text = str(sample.get("output") or "")
parts = [p for p in [instruction, input_text, output_text] if p]
return "\n".join(parts)
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=str, required=True)
args = parser.parse_args()
cfg = load_config(args.config)
os.environ["TOKENIZERS_PARALLELISM"] = "false"
tokenizer = AutoTokenizer.from_pretrained(cfg.model_id, use_fast=True)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
model = AutoModelForCausalLM.from_pretrained(cfg.model_id)
lora_config = LoraConfig(
task_type=TaskType.CAUSAL_LM,
r=cfg.lora_r,
lora_alpha=cfg.lora_alpha,
lora_dropout=cfg.lora_dropout,
)
model = get_peft_model(model, lora_config)
raw_dataset = load_dataset("json", data_files={"train": cfg.dataset_path})
def tokenize_fn(batch: Dict[str, List[str]]) -> Dict[str, Any]:
texts = [format_sample({"text": t}) for t in batch["text"]] if "text" in batch else [
format_sample({k: v[i] for k, v in batch.items()}) for i in range(len(next(iter(batch.values()))))
]
return tokenizer(texts, truncation=True, max_length=cfg.max_seq_length)
tokenized = raw_dataset.map(tokenize_fn, batched=True, remove_columns=raw_dataset["train"].column_names)
training_args = TrainingArguments(
output_dir=cfg.output_dir,
num_train_epochs=cfg.num_train_epochs,
per_device_train_batch_size=cfg.per_device_train_batch_size,
gradient_accumulation_steps=cfg.gradient_accumulation_steps,
learning_rate=cfg.learning_rate,
logging_steps=10,
save_steps=50,
save_total_limit=2,
seed=cfg.seed,
fp16=False,
report_to=[],
)
data_collator = DataCollatorForLanguageModeling(tokenizer=tokenizer, mlm=False)
trainer = Trainer(
model=model,
args=training_args,
train_dataset=tokenized["train"],
data_collator=data_collator,
)
trainer.train()
trainer.save_model(cfg.output_dir)
tokenizer.save_pretrained(cfg.output_dir)
print(f"Training complete. Saved to {cfg.output_dir}")
if __name__ == "__main__":
main()