Ship Assistent 0.13.0: real QLoRA pipeline and GGUF Ollama register.
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>
This commit is contained in:
+198
-38
@@ -1,12 +1,16 @@
|
||||
#!/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."""
|
||||
"""QLoRA SFT trainer for Swarm Assistent."""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
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):
|
||||
@@ -17,70 +21,226 @@ def log(msg, log_path):
|
||||
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)
|
||||
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)
|
||||
|
||||
log(f"Swarm Assistent QLoRA starting base={cfg.get('base_model')}", args.log)
|
||||
|
||||
try:
|
||||
import torch
|
||||
from datasets import load_dataset
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer, TrainingArguments, Trainer
|
||||
from peft import LoraConfig, get_peft_model, TaskType
|
||||
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 torch transformers datasets peft bitsandbytes trl accelerate", args.log)
|
||||
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,
|
||||
load_in_4bit=bool(cfg.get("four_bit", True)),
|
||||
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=int(cfg.get("rank", 16)),
|
||||
lora_alpha=int(cfg.get("alpha", 32)),
|
||||
r=rank,
|
||||
lora_alpha=alpha,
|
||||
lora_dropout=0.05,
|
||||
task_type=TaskType.CAUSAL_LM,
|
||||
target_modules=["q_proj", "v_proj", "k_proj", "o_proj"],
|
||||
target_modules=target_modules,
|
||||
)
|
||||
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}
|
||||
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)
|
||||
|
||||
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)
|
||||
json.dump({"ok": True, "base": base, "rows": len(ds)}, f)
|
||||
|
||||
log("training complete", args.log)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user