#!/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()