captchaforge 0.2.28

Automatic CAPTCHA detection and multi-strategy solving for chromiumoxide-driven headless browsers (Cloudflare Turnstile, reCAPTCHA v2/v3, hCaptcha, image grids, audio, sliders).
Documentation
#!/usr/bin/env python3
"""LoRA fine-tune a vision-language model on a captchaforge VlmDataset.

Reads the manifest.jsonl produced by `captchaforge::vlm_dataset`,
loads each (image, prompt, answer) sample into the format the
chosen base model expects, runs a LoRA fine-tune, writes the
adapter weights + a merged-checkpoint to the output dir.

Defaults target Qwen2-VL-7B because (a) it's open-weights, (b)
its single-image inference cost is acceptable for production
captcha solving, (c) the base model already does decent
image-grid + text reading out of the box, so the fine-tune just
needs to teach it captcha-specific output formatting.

Usage:
    train_lora.py --base-model qwen/qwen2-vl-7b \\
                  --dataset /work/datasets/round-N \\
                  --out /work/weights/captchaforge-vlm-round-N \\
                  --epochs 3 --batch-size 4
"""
from __future__ import annotations

import argparse
import dataclasses
import json
import logging
import os
import sys
from pathlib import Path
from typing import Iterable

# Heavy ML imports are deferred until after argparse so `--help`
# is fast and dry-runs (e.g. CI smoke-tests) don't pay the cost.

LOG = logging.getLogger("captchaforge.train_lora")


@dataclasses.dataclass
class VlmSample:
    """One sample as serialised by `captchaforge::vlm_dataset`."""

    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]:
    """Stream the JSONL manifest into a list of VlmSamples."""
    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:
    """Convert a captchaforge sample to the HF chat format Qwen2-VL expects."""
    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

    # Imports deferred until we're actually about to train.
    LOG.info("importing torch / transformers / peft …")
    import torch  # noqa: WPS433
    from transformers import AutoProcessor, Qwen2VLForConditionalGeneration  # noqa: WPS433
    from peft import LoraConfig, get_peft_model, TaskType  # noqa: WPS433
    from trl import SFTConfig, SFTTrainer  # noqa: WPS433
    from datasets import Dataset  # noqa: WPS433

    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,
        # Qwen2-VL standard target modules.
        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=[],  # silence wandb/etc; ops can opt in via env
        dataset_text_field=None,  # multi-modal — not text-only
        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:]))