bae705aa97
- Add NFC ePassport roadmap (ICAO 9303, eIDAS) - Add TensorFlow.js edge face detection (BlazeFace) - Add structured audit logger (GDPR-compliant) - Risk scoring support Part of KYC Apple Native UX v1.1.0
231 lines
9.8 KiB
Python
231 lines
9.8 KiB
Python
#!/usr/bin/env python3
|
|
"""
|
|
train_eslm.py — Rekonstruerat träningsscript för ESLM / amos-r2
|
|
================================================================
|
|
Status: REKONSTRUERAT (originalscriptet hittades ej på GPU-boxen)
|
|
|
|
Evidensbas för hyperparametrar:
|
|
- adapter_config.json → lora_r=16, lora_alpha=32, lora_dropout=0.05,
|
|
target_modules (alla 7 proj), use_dora=False,
|
|
use_rslora=False
|
|
- UnslothSFTConfig defaults (train_env/unsloth 2026.5.8, trl 0.24.0):
|
|
per_device_train_batch_size=4, gradient_accumulation_steps=2,
|
|
learning_rate=5e-5, num_train_epochs=3, optim=adamw_8bit,
|
|
warmup_steps=0.1, weight_decay=0.001, seed=3407
|
|
- tokenizer_config.json → padding_side=left, eos=<|im_end|>,
|
|
model_max_length=32768
|
|
- config.json (aamos_merged) → 4bit NF4 + bfloat16 compute,
|
|
double_quant=True
|
|
- PEFT version 0.19.1 → init_lora_weights=True (standard)
|
|
- Qwen2.5 chat template → assistant_only_loss=True (sannolikt)
|
|
|
|
Osäkra parametrar (markeras med "GUESS"):
|
|
- max_seq_length: 2048 (GUESS — vllm-serve kör max_model_len=4096 men
|
|
träning brukar halveras; 2048 är vanligt för Qwen2.5 SFT-demos)
|
|
- num_train_epochs: 3 (Unsloth default; okänd)
|
|
- Dataset-path: okänd
|
|
|
|
Kör INTE träning direkt — verifiera att --data och --output är rätt.
|
|
"""
|
|
|
|
import argparse
|
|
import json
|
|
import os
|
|
from pathlib import Path
|
|
|
|
def parse_args():
|
|
parser = argparse.ArgumentParser(
|
|
description="Träna Qwen2.5-7B LoRA via Unsloth (amos-r2 rekonstruktion)"
|
|
)
|
|
parser.add_argument(
|
|
"--data", required=True,
|
|
help="Sökväg till JSONL-fil med messages-format: "
|
|
'[{"role":"system","content":"..."},{"role":"user","content":"..."},{"role":"assistant","content":"..."}]'
|
|
)
|
|
parser.add_argument(
|
|
"--output", required=True,
|
|
help="Output-katalog för LoRA-adaptern"
|
|
)
|
|
parser.add_argument(
|
|
"--base-model", default="unsloth/Qwen2.5-7B-Instruct-bnb-4bit",
|
|
help="Bas-modell (default: unsloth/Qwen2.5-7B-Instruct-bnb-4bit)"
|
|
)
|
|
parser.add_argument(
|
|
"--max-seq-length", type=int, default=2048,
|
|
help="Max sekvenslängd (GUESS: 2048; justera vid OOM)"
|
|
)
|
|
parser.add_argument(
|
|
"--epochs", type=float, default=3.0,
|
|
help="Antal träningsepoks (Unsloth default: 3)"
|
|
)
|
|
parser.add_argument(
|
|
"--lr", type=float, default=5e-5,
|
|
help="Learning rate (Unsloth default: 5e-5)"
|
|
)
|
|
parser.add_argument(
|
|
"--batch-size", type=int, default=4,
|
|
help="Per-device train batch size (Unsloth default: 4)"
|
|
)
|
|
parser.add_argument(
|
|
"--grad-accum", type=int, default=2,
|
|
help="Gradient accumulation steps (Unsloth default: 2)"
|
|
)
|
|
parser.add_argument(
|
|
"--merge-output", default=None,
|
|
help="Om angiven, merga adapter+bas till denna katalog"
|
|
)
|
|
return parser.parse_args()
|
|
|
|
|
|
def load_dataset_from_jsonl(path: str):
|
|
"""Ladda JSONL med messages-format till HuggingFace Dataset."""
|
|
from datasets import Dataset
|
|
records = []
|
|
with open(path, "r", encoding="utf-8") as f:
|
|
for line in f:
|
|
line = line.strip()
|
|
if not line:
|
|
continue
|
|
obj = json.loads(line)
|
|
# Förväntat format: {"messages": [...]} eller direkt lista
|
|
if isinstance(obj, list):
|
|
records.append({"messages": obj})
|
|
elif "messages" in obj:
|
|
records.append({"messages": obj["messages"]})
|
|
else:
|
|
# Försök tolka som system/user/assistant-keys
|
|
msgs = []
|
|
if "system" in obj:
|
|
msgs.append({"role": "system", "content": obj["system"]})
|
|
if "user" in obj or "input" in obj:
|
|
msgs.append({"role": "user", "content": obj.get("user", obj.get("input", ""))})
|
|
if "assistant" in obj or "output" in obj:
|
|
msgs.append({"role": "assistant", "content": obj.get("assistant", obj.get("output", ""))})
|
|
if msgs:
|
|
records.append({"messages": msgs})
|
|
else:
|
|
raise ValueError(f"Okänt dataformat i rad: {line[:100]}")
|
|
return Dataset.from_list(records)
|
|
|
|
|
|
def main():
|
|
args = parse_args()
|
|
|
|
print(f"[train_eslm] Laddar Unsloth + modell...")
|
|
from unsloth import FastLanguageModel
|
|
from trl import SFTTrainer
|
|
from unsloth.trainer import UnslothTrainingArguments
|
|
|
|
# ─── 1. Ladda bas-modell (4-bit NF4 + bfloat16) ───────────────────────────
|
|
model, tokenizer = FastLanguageModel.from_pretrained(
|
|
model_name=args.base_model,
|
|
max_seq_length=args.max_seq_length,
|
|
dtype=None, # auto (bfloat16 om GPU stöder det)
|
|
load_in_4bit=True,
|
|
)
|
|
|
|
# ─── 2. Applicera LoRA (parametrar från adapter_config.json) ──────────────
|
|
model = FastLanguageModel.get_peft_model(
|
|
model,
|
|
r=16, # från adapter_config.json: "r": 16
|
|
lora_alpha=32, # från adapter_config.json: "lora_alpha": 32
|
|
lora_dropout=0.05, # från adapter_config.json: "lora_dropout": 0.05
|
|
target_modules=[ # från adapter_config.json: "target_modules"
|
|
"q_proj", "k_proj", "v_proj", "o_proj",
|
|
"gate_proj", "up_proj", "down_proj",
|
|
],
|
|
bias="none", # från adapter_config.json: "bias": "none"
|
|
use_gradient_checkpointing="unsloth",
|
|
random_state=3407, # Unsloth standard seed
|
|
use_rslora=False, # från adapter_config.json: "use_rslora": false
|
|
use_dora=False, # från adapter_config.json: "use_dora": false
|
|
loftq_config=None,
|
|
)
|
|
|
|
# ─── 3. Ladda dataset ─────────────────────────────────────────────────────
|
|
print(f"[train_eslm] Laddar dataset från {args.data}...")
|
|
dataset = load_dataset_from_jsonl(args.data)
|
|
print(f"[train_eslm] Dataset: {len(dataset)} exempel")
|
|
|
|
# ─── 4. Förbered chat-format (Qwen2.5 Instruct-mall) ─────────────────────
|
|
# Qwen2.5-Instruct använder ChatML: <|im_start|>role\ncontent<|im_end|>
|
|
def format_example(example):
|
|
text = tokenizer.apply_chat_template(
|
|
example["messages"],
|
|
tokenize=False,
|
|
add_generation_prompt=False,
|
|
)
|
|
return {"text": text}
|
|
|
|
dataset = dataset.map(format_example, desc="Formaterar med chat template")
|
|
|
|
# ─── 5. Träningskonfiguration (Unsloth defaults + adapter-evidens) ────────
|
|
# Obs: UnslothTrainingArguments wraps TrainingArguments med Unsloth-optimeringar
|
|
# Fallback till standard SFTConfig om UnslothTrainingArguments ej finns
|
|
try:
|
|
from unsloth.trainer import UnslothTrainingArguments as TrainingArgs
|
|
except ImportError:
|
|
try:
|
|
from unsloth_compiled_cache.UnslothSFTTrainer import UnslothSFTConfig as TrainingArgs
|
|
except ImportError:
|
|
from trl import SFTConfig as TrainingArgs
|
|
|
|
training_args = TrainingArgs(
|
|
output_dir=args.output,
|
|
per_device_train_batch_size=args.batch_size, # 4 (Unsloth default)
|
|
gradient_accumulation_steps=args.grad_accum, # 2 (Unsloth default)
|
|
num_train_epochs=args.epochs, # 3 (Unsloth default)
|
|
learning_rate=args.lr, # 5e-5 (Unsloth default)
|
|
lr_scheduler_type="linear", # Unsloth default
|
|
warmup_steps=0.1, # Unsloth default (10% av steg)
|
|
optim="adamw_8bit", # Unsloth default
|
|
weight_decay=0.001, # Unsloth default
|
|
adam_beta1=0.9,
|
|
adam_beta2=0.999,
|
|
max_grad_norm=1.0,
|
|
bf16=True, # bfloat16 (från quantization_config)
|
|
fp16=False,
|
|
gradient_checkpointing=True,
|
|
seed=3407, # Unsloth standard
|
|
logging_steps=1,
|
|
logging_strategy="steps",
|
|
save_strategy="steps",
|
|
save_steps=500,
|
|
report_to="none",
|
|
dataset_text_field="text",
|
|
max_seq_length=args.max_seq_length, # 2048 (GUESS)
|
|
)
|
|
|
|
# ─── 6. Kör SFT-träning ───────────────────────────────────────────────────
|
|
trainer = SFTTrainer(
|
|
model=model,
|
|
tokenizer=tokenizer,
|
|
args=training_args,
|
|
train_dataset=dataset,
|
|
)
|
|
|
|
print("[train_eslm] Startar träning...")
|
|
trainer_stats = trainer.train()
|
|
print(f"[train_eslm] Träning klar: {trainer_stats}")
|
|
|
|
# ─── 7. Spara LoRA-adaptern ───────────────────────────────────────────────
|
|
print(f"[train_eslm] Sparar LoRA-adapter till {args.output}...")
|
|
model.save_pretrained(args.output)
|
|
tokenizer.save_pretrained(args.output)
|
|
|
|
# ─── 8. (Optionellt) Mergea adapter med bas-modell ────────────────────────
|
|
if args.merge_output:
|
|
print(f"[train_eslm] Mergear modell till {args.merge_output}...")
|
|
model.save_pretrained_merged(
|
|
args.merge_output,
|
|
tokenizer,
|
|
save_method="merged_16bit", # spara som float16 för vLLM
|
|
)
|
|
print(f"[train_eslm] Merged modell sparad till {args.merge_output}")
|
|
|
|
print("[train_eslm] Klart!")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|