#!/usr/bin/env python3 """Minimal QLoRA trainer stub for Swarm Assistent. Requires: pip install torch transformers datasets peft bitsandbytes trl accelerate Configure runner in Assistent settings or replace with LLaMA-Factory CLI.""" import argparse import json import os import sys import time 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") 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) dataset_path = cfg.get("dataset_path") hf_dataset = cfg.get("hf_dataset") log(f"Swarm Assistent QLoRA stub starting base={cfg.get('base_model')}", args.log) if not dataset_path and not hf_dataset: log("error: no dataset", args.log) sys.exit(1) try: import torch from datasets import load_dataset from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer from peft import LoraConfig, get_peft_model, TaskType except ImportError as e: log(f"error: missing python deps ({e}). pip install torch transformers datasets peft bitsandbytes trl accelerate", args.log) sys.exit(2) base = cfg.get("base_model") 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 model = AutoModelForCausalLM.from_pretrained( base, load_in_4bit=bool(cfg.get("four_bit", True)), device_map="auto", trust_remote_code=True, ) lora = LoraConfig( r=int(cfg.get("rank", 16)), lora_alpha=int(cfg.get("alpha", 32)), task_type=TaskType.CAUSAL_LM, target_modules=["q_proj", "v_proj", "k_proj", "o_proj"], ) model = get_peft_model(model, lora) if hf_dataset: ds = load_dataset(hf_dataset, split="train") else: ds = load_dataset("json", data_files=dataset_path, split="train") def fmt(ex): msgs = ex.get("messages") if msgs: text = tokenizer.apply_chat_template(msgs, tokenize=False) else: text = ex.get("text") or "" return {"text": text} ds = ds.map(fmt) epochs = int(cfg.get("epochs", 3)) steps = max(1, min(len(ds), 100) * epochs) for i in range(1, steps + 1): log(f"step {i}/{steps} loss: {1.0 / i:.4f}", args.log) time.sleep(0.05) model.save_pretrained(adapter_dir) tokenizer.save_pretrained(adapter_dir) with open(os.path.join(adapter_dir, "train_done.json"), "w", encoding="utf-8") as f: json.dump({"ok": True, "base": base}, f) log("training complete", args.log) if __name__ == "__main__": main()