from __future__ import annotations
import argparse
import dataclasses
import json
import logging
import os
import sys
from pathlib import Path
from typing import Iterable
LOG = logging.getLogger("captchaforge.train_lora")
@dataclasses.dataclass
class VlmSample:
id: str
kind: str
image_path: str
prompt: str
answer: str
is_negative: bool
def parse_args(argv: list[str]) -> argparse.Namespace:
p = argparse.ArgumentParser(
prog="train_lora.py",
description="LoRA fine-tune a VLM on a captchaforge VlmDataset.",
formatter_class=argparse.ArgumentDefaultsHelpFormatter,
)
p.add_argument(
"--base-model",
default="Qwen/Qwen2-VL-7B-Instruct",
help="HuggingFace base model. Must be vision-capable.",
)
p.add_argument(
"--dataset",
type=Path,
required=True,
help="Path to a captchaforge VlmDataset directory (containing manifest.jsonl + images/).",
)
p.add_argument(
"--out",
type=Path,
required=True,
help="Output dir for the LoRA adapter + merged-checkpoint.",
)
p.add_argument("--epochs", type=int, default=3)
p.add_argument("--batch-size", type=int, default=4)
p.add_argument(
"--lr",
type=float,
default=2e-4,
help="LoRA peak learning rate. PEFT defaults to 1e-4; we run hotter for captcha.",
)
p.add_argument(
"--lora-r",
type=int,
default=16,
help="LoRA rank. 8-32 is the sensible band; higher = more capacity, slower.",
)
p.add_argument(
"--lora-alpha",
type=int,
default=32,
help="LoRA scaling factor. Convention: 2× rank.",
)
p.add_argument(
"--max-samples",
type=int,
default=None,
help="Cap dataset to N samples (useful for smoke-tests). None = use all.",
)
p.add_argument(
"--dry-run",
action="store_true",
help="Parse the dataset + print stats; skip model load + training.",
)
p.add_argument(
"--log-level",
default="INFO",
choices=["DEBUG", "INFO", "WARNING", "ERROR"],
)
return p.parse_args(argv)
def load_manifest(dataset_dir: Path, max_samples: int | None = None) -> list[VlmSample]:
manifest_path = dataset_dir / "manifest.jsonl"
if not manifest_path.exists():
raise FileNotFoundError(f"manifest not found at {manifest_path}")
samples: list[VlmSample] = []
with manifest_path.open("r", encoding="utf-8") as f:
for line_no, raw in enumerate(f, start=1):
raw = raw.strip()
if not raw:
continue
try:
obj = json.loads(raw)
except json.JSONDecodeError as e:
LOG.warning("skipping malformed manifest line %d: %s", line_no, e)
continue
samples.append(
VlmSample(
id=obj.get("id", ""),
kind=obj.get("kind", ""),
image_path=obj["image_path"],
prompt=obj.get("prompt", ""),
answer=obj.get("answer", ""),
is_negative=obj.get("is_negative", False),
)
)
if max_samples is not None and len(samples) >= max_samples:
break
return samples
def per_kind_counts(samples: Iterable[VlmSample]) -> dict[str, int]:
out: dict[str, int] = {}
for s in samples:
out[s.kind] = out.get(s.kind, 0) + 1
return out
def to_chat_format(sample: VlmSample, dataset_dir: Path) -> dict:
image_full_path = (dataset_dir / sample.image_path).as_posix()
return {
"messages": [
{
"role": "user",
"content": [
{"type": "image", "image": image_full_path},
{"type": "text", "text": sample.prompt},
],
},
{
"role": "assistant",
"content": [{"type": "text", "text": sample.answer}],
},
]
}
def main(argv: list[str]) -> int:
args = parse_args(argv)
logging.basicConfig(
level=getattr(logging, args.log_level),
format="%(asctime)s %(levelname)s %(name)s %(message)s",
)
LOG.info("loading manifest from %s", args.dataset)
samples = load_manifest(args.dataset, max_samples=args.max_samples)
LOG.info("loaded %d samples", len(samples))
counts = per_kind_counts(samples)
for kind, n in sorted(counts.items()):
LOG.info(" %4d %s", n, kind)
if not samples:
LOG.error("dataset is empty — nothing to train on")
return 2
if args.dry_run:
LOG.info("dry-run mode; skipping model load")
return 0
LOG.info("importing torch / transformers / peft …")
import torch from transformers import AutoProcessor, Qwen2VLForConditionalGeneration from peft import LoraConfig, get_peft_model, TaskType from trl import SFTConfig, SFTTrainer from datasets import Dataset
LOG.info("loading base model %s …", args.base_model)
processor = AutoProcessor.from_pretrained(args.base_model)
model = Qwen2VLForConditionalGeneration.from_pretrained(
args.base_model,
torch_dtype=torch.bfloat16,
device_map="auto",
)
lora_cfg = LoraConfig(
r=args.lora_r,
lora_alpha=args.lora_alpha,
target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
lora_dropout=0.05,
bias="none",
task_type=TaskType.CAUSAL_LM,
)
model = get_peft_model(model, lora_cfg)
model.print_trainable_parameters()
LOG.info("converting %d samples to HF chat format …", len(samples))
chat_samples = [to_chat_format(s, args.dataset) for s in samples]
ds = Dataset.from_list(chat_samples)
args.out.mkdir(parents=True, exist_ok=True)
sft_args = SFTConfig(
output_dir=str(args.out),
num_train_epochs=args.epochs,
per_device_train_batch_size=args.batch_size,
gradient_accumulation_steps=4,
learning_rate=args.lr,
bf16=True,
logging_steps=10,
save_strategy="epoch",
report_to=[], dataset_text_field=None, max_seq_length=2048,
)
trainer = SFTTrainer(
model=model,
args=sft_args,
train_dataset=ds,
processing_class=processor,
)
LOG.info("starting fine-tune …")
trainer.train()
LOG.info("saving adapter + merged checkpoint to %s …", args.out)
trainer.save_model(str(args.out))
processor.save_pretrained(str(args.out))
LOG.info("done. promote with scripts/promote_to_ollama.sh %s", args.out)
return 0
if __name__ == "__main__":
sys.exit(main(sys.argv[1:]))