Replace fake train loop with TRL SFTTrainer, HF column mapping with fiction preset, safetensors to GGUF conversion, and ollama create using ollama_base plus ADAPTER. Co-authored-by: Cursor <cursoragent@cursor.com>
249 lines
7.8 KiB
Python
249 lines
7.8 KiB
Python
#!/usr/bin/env python3
|
|
"""QLoRA SFT trainer for Swarm Assistent."""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import os
|
|
import sys
|
|
import traceback
|
|
|
|
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
|
|
|
from hf_dataset_map import row_to_messages
|
|
|
|
|
|
def log(msg, log_path):
|
|
line = str(msg)
|
|
print(line, flush=True)
|
|
if log_path:
|
|
with open(log_path, "a", encoding="utf-8") as f:
|
|
f.write(line + "\n")
|
|
|
|
|
|
class StepLogger:
|
|
def __init__(self, log_path, total_steps):
|
|
self.log_path = log_path
|
|
self.total_steps = max(1, total_steps)
|
|
|
|
def on_log(self, logs):
|
|
loss = logs.get("loss")
|
|
step = logs.get("step") or logs.get("global_step")
|
|
if step is None:
|
|
return
|
|
if loss is not None:
|
|
log(f"step {int(step)}/{self.total_steps} loss: {float(loss):.4f}", self.log_path)
|
|
|
|
|
|
def detect_target_modules(model):
|
|
names = {n.split(".")[-1] for n, _ in model.named_modules()}
|
|
candidates = [
|
|
["q_proj", "k_proj", "v_proj", "o_proj"],
|
|
["q_proj", "v_proj"],
|
|
["Wqkv", "out_proj"],
|
|
["c_attn", "c_proj"],
|
|
]
|
|
for group in candidates:
|
|
if all(g in names for g in group):
|
|
return group
|
|
return ["q_proj", "v_proj"]
|
|
|
|
|
|
def load_sft_dataset(cfg, log_path):
|
|
from datasets import Dataset, load_dataset
|
|
|
|
hf_dataset = cfg.get("hf_dataset")
|
|
dataset_path = cfg.get("dataset_path")
|
|
max_samples = int(cfg.get("max_samples") or 0)
|
|
hf_mapping = cfg.get("hf_mapping") or {}
|
|
schema = cfg.get("hf_schema") or {}
|
|
|
|
rows = []
|
|
|
|
if hf_dataset:
|
|
log(f"loading HF dataset {hf_dataset}", log_path)
|
|
ds = load_dataset(hf_dataset, split="train")
|
|
if max_samples > 0:
|
|
ds = ds.select(range(min(len(ds), max_samples)))
|
|
for ex in ds:
|
|
msgs = row_to_messages(dict(ex), schema, hf_mapping)
|
|
if msgs:
|
|
rows.append({"messages": msgs})
|
|
elif dataset_path and os.path.isfile(dataset_path):
|
|
log(f"loading JSONL {dataset_path}", log_path)
|
|
with open(dataset_path, encoding="utf-8") as f:
|
|
for line in f:
|
|
line = line.strip()
|
|
if not line:
|
|
continue
|
|
try:
|
|
obj = json.loads(line)
|
|
except json.JSONDecodeError:
|
|
continue
|
|
msgs = obj.get("messages")
|
|
if msgs:
|
|
rows.append({"messages": msgs})
|
|
if max_samples > 0 and len(rows) >= max_samples:
|
|
break
|
|
else:
|
|
log("error: no dataset_path or hf_dataset", log_path)
|
|
sys.exit(1)
|
|
|
|
if not rows:
|
|
log("error: no training rows after mapping", log_path)
|
|
sys.exit(1)
|
|
|
|
log(f"dataset rows: {len(rows)}", log_path)
|
|
return Dataset.from_list(rows)
|
|
|
|
|
|
def main():
|
|
p = argparse.ArgumentParser()
|
|
p.add_argument("--config", required=True)
|
|
p.add_argument("--log", required=True)
|
|
args = p.parse_args()
|
|
|
|
with open(args.config, encoding="utf-8") as f:
|
|
cfg = json.load(f)
|
|
|
|
adapter_dir = cfg.get("adapter_dir", "adapter")
|
|
os.makedirs(adapter_dir, exist_ok=True)
|
|
|
|
log(f"Swarm Assistent QLoRA starting base={cfg.get('base_model')}", args.log)
|
|
|
|
try:
|
|
import torch
|
|
from peft import LoraConfig, TaskType, get_peft_model
|
|
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig, TrainerCallback, TrainingArguments
|
|
from trl import SFTTrainer
|
|
except ImportError as e:
|
|
log(f"error: missing python deps ({e}). pip install -r scripts/requirements-train.txt", args.log)
|
|
sys.exit(2)
|
|
|
|
base = cfg.get("base_model")
|
|
if not base:
|
|
log("error: base_model required", args.log)
|
|
sys.exit(1)
|
|
|
|
try:
|
|
ds = load_sft_dataset(cfg, args.log)
|
|
except Exception as e:
|
|
log(f"error: dataset load failed: {e}", args.log)
|
|
traceback.print_exc()
|
|
sys.exit(1)
|
|
|
|
four_bit = bool(cfg.get("four_bit", True))
|
|
seq_len = int(cfg.get("seq_len") or 2048)
|
|
rank = int(cfg.get("rank") or 16)
|
|
alpha = int(cfg.get("alpha") or 32)
|
|
lr = float(cfg.get("lr") or 2e-4)
|
|
epochs = int(cfg.get("epochs") or 3)
|
|
batch_size = int(cfg.get("batch_size") or 1)
|
|
grad_accum = int(cfg.get("gradient_accumulation_steps") or 4)
|
|
|
|
log(f"loading model {base}", args.log)
|
|
tokenizer = AutoTokenizer.from_pretrained(base, trust_remote_code=True)
|
|
if tokenizer.pad_token is None:
|
|
tokenizer.pad_token = tokenizer.eos_token
|
|
|
|
bnb_config = None
|
|
if four_bit:
|
|
bnb_config = BitsAndBytesConfig(
|
|
load_in_4bit=True,
|
|
bnb_4bit_quant_type="nf4",
|
|
bnb_4bit_compute_dtype=torch.float16,
|
|
bnb_4bit_use_double_quant=True,
|
|
)
|
|
|
|
model = AutoModelForCausalLM.from_pretrained(
|
|
base,
|
|
quantization_config=bnb_config,
|
|
device_map="auto",
|
|
trust_remote_code=True,
|
|
torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
|
|
)
|
|
|
|
target_modules = detect_target_modules(model)
|
|
log(f"lora target_modules: {target_modules}", args.log)
|
|
|
|
lora = LoraConfig(
|
|
r=rank,
|
|
lora_alpha=alpha,
|
|
lora_dropout=0.05,
|
|
task_type=TaskType.CAUSAL_LM,
|
|
target_modules=target_modules,
|
|
)
|
|
model = get_peft_model(model, lora)
|
|
|
|
def formatting_func(examples):
|
|
texts = []
|
|
for msgs in examples["messages"]:
|
|
try:
|
|
text = tokenizer.apply_chat_template(msgs, tokenize=False, add_generation_prompt=False)
|
|
except Exception:
|
|
parts = []
|
|
for m in msgs:
|
|
role = m.get("role", "user")
|
|
content = m.get("content", "")
|
|
parts.append(f"{role}: {content}")
|
|
text = "\n".join(parts)
|
|
texts.append(text)
|
|
return texts
|
|
|
|
total_steps = max(1, (len(ds) * epochs) // max(1, batch_size * grad_accum))
|
|
log(f"planned steps ~{total_steps}", args.log)
|
|
|
|
training_args = TrainingArguments(
|
|
output_dir=adapter_dir,
|
|
num_train_epochs=epochs,
|
|
per_device_train_batch_size=batch_size,
|
|
gradient_accumulation_steps=grad_accum,
|
|
learning_rate=lr,
|
|
logging_steps=1,
|
|
save_steps=max(50, total_steps // 10),
|
|
save_total_limit=2,
|
|
fp16=torch.cuda.is_available(),
|
|
bf16=False,
|
|
report_to="none",
|
|
remove_unused_columns=False,
|
|
max_grad_norm=0.3,
|
|
warmup_ratio=0.03,
|
|
lr_scheduler_type="cosine",
|
|
)
|
|
|
|
step_logger = StepLogger(args.log, total_steps)
|
|
|
|
class LossCallback(TrainerCallback):
|
|
def on_log(self, args_, state, control, logs=None, **kwargs):
|
|
if logs:
|
|
step_logger.on_log({**logs, "step": state.global_step})
|
|
|
|
try:
|
|
trainer = SFTTrainer(
|
|
model=model,
|
|
args=training_args,
|
|
train_dataset=ds,
|
|
tokenizer=tokenizer,
|
|
formatting_func=formatting_func,
|
|
max_seq_length=seq_len,
|
|
packing=False,
|
|
callbacks=[LossCallback()],
|
|
)
|
|
log("training started", args.log)
|
|
trainer.train()
|
|
trainer.save_model(adapter_dir)
|
|
tokenizer.save_pretrained(adapter_dir)
|
|
except Exception as e:
|
|
log(f"error: training failed: {e}", args.log)
|
|
traceback.print_exc()
|
|
sys.exit(3)
|
|
|
|
with open(os.path.join(adapter_dir, "train_done.json"), "w", encoding="utf-8") as f:
|
|
json.dump({"ok": True, "base": base, "rows": len(ds)}, f)
|
|
|
|
log("training complete", args.log)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|