feat(trainer): Optimizer & PEFT Zoo (v0.41.0)

Optimizer Zoo (Part A): closed-allowlist SUPPORTED_OPTIMIZERS adds 14 new
entries — BAdam, APOLLO (apollo_adamw), Adam-mini, lomo / adalomo,
grokadamw, schedule_free_adamw / schedule_free_sgd, muon, dion,
came_pytorch, ao_adamw_{fp8,4bit,8bit}. Validates name + lower-cases
deterministically; rejects non-string / empty / null-byte / >64-char.
_OPTIMIZER_PACKAGES wrapped in MappingProxyType (matches v0.36.0 _REGISTRY).

Per-module LR (Part B): training.lr_groups accepts list-of-pairs /
list-of-dicts / {pattern: lr} mapping; canonical [{pattern, lr}, ...]
storage. Capped at MAX_LR_GROUPS=32; per-pattern non-empty string
≤256 chars + null-byte rejection + re.compile + best-effort ReDoS probe;
per-LR (0.0, 1.0] + math.isfinite (rejects NaN AND ±inf) + bool rejection
(matches v0.30.0 Candidate policy); duplicates rejected. lr_groups_from_schema
bridges canonical schema shape into runtime List[LrGroup] for
build_optimizer_param_groups (first-match-wins routing). LrGroup is
@dataclass(frozen=True). PyYAML scientific-notation (1e-4) parses as
string in YAML 1.1; _validate_lr coerces str → float so soup.yaml
round-trips work. base_lr rejects bool / non-positive (defence-in-depth).

PEFT methods (Part C): LoraConfig.init_strategy="loftq" + loftq_iter
∈ [1, 10] + loftq_bits ∈ {2, 4, 8}; cross-validator rejects loftq +
use_dora / use_vera. utils/loftq_init.py exposes validators +
build_loftq_config (lazy peft.LoftQConfig with actionable ImportError
hint). LLaMA Pro: TrainingConfig.expand_layers ∈ [1, 64] +
freeze_trainable_layers (signed, |x| ≤ 1000); cross-validator requires
the pair (LLaMA Pro freezes original layers and trains only new blocks).
field_validator(mode="before") on both rejects bool BEFORE Pydantic ge/le
silently coerces True → 1. expand_model_blocks raises NotImplementedError
with v0.41.1 marker — schema-only release (mirrors v0.27.0 / v0.37.0
stub-then-live pattern). utils/block_expansion._count_layers uses
hasattr(__len__) instead of try/except TypeError so legitimate __len__
bugs surface loudly. use_mod boolean for Mixture-of-Depths (schema only —
live patch deferred to v0.41.1).

Friendly aliases: load_in_8bit / load_in_16bit (Optional[bool]) for
LlamaFactory / Axolotl users. is True policy on both — explicit False is
"no preference", not "off"; mutually-exclusive both-True rejected; alias
combined with explicit Quant Menu format raises rather than silently
overriding. Alias-driven quantization rewrite via direct
self.quantization = ... (Pydantic v2 BaseModel non-frozen path), NOT
object.__setattr__ — code review caught that the latter would silently
bypass any future field_validator on quantization.

Five review agents (4 in parallel + manual smoke for verification-loop):
all CRITICAL/HIGH/MEDIUM/LOW findings fixed. Local smoke caught a real
bug — PyYAML parsing 1e-4 as string — fixed pre-commit with 2 added tests.

5242 tests pass (+120 net new); ruff clean.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Alpamys 2026-05-10 15:35:35 +05:00
parent 5e0872b9ea
commit 1491025a36
14 changed files with 1454 additions and 16 deletions

View File

@ -107,11 +107,11 @@ soup_cli/
cans/ - Shareable .can artifact format + run/publish orchestrator (v0.26.0 + v0.33.0)
data/traces/ - Trace-to-Preference harvester (v0.26.0)
data/collators.py - CrossDocCollator for sample packing (v0.33.0)
utils/ - GPU, errors, MoE, GaLore, QAT, Unsloth, vLLM, SGLang, Liger, FlashAttn, FSDP, Ring Attention, long-context, quality, curriculum, freeze, dataset-registry, mlx, peft_builder, paths, topology, launcher, mii, pipeline, cut_ce, fp8, gradient_ckpt, kernel_picker, cross_doc_attn, activation_offload, hf, spec_pairing, structured_output, metrics, tracing, auto_quant, lr_finder, grad_accum, mixed_precision, warmup, spike_recovery, convergence, v028_features, multipack_sampler, multipack, neat_packing, jinja_analyzer, quant_menu, relora, peft_patches, peft_wiring, dpo_variants
utils/ - GPU, errors, MoE, GaLore, QAT, Unsloth, vLLM, SGLang, Liger, FlashAttn, FSDP, Ring Attention, long-context, quality, curriculum, freeze, dataset-registry, mlx, peft_builder, paths, topology, launcher, mii, pipeline, cut_ce, fp8, gradient_ckpt, kernel_picker, cross_doc_attn, activation_offload, hf, spec_pairing, structured_output, metrics, tracing, auto_quant, lr_finder, grad_accum, mixed_precision, warmup, spike_recovery, convergence, v028_features, multipack_sampler, multipack, neat_packing, jinja_analyzer, quant_menu, relora, peft_patches, peft_wiring, dpo_variants, optimizer_zoo, lr_groups, loftq_init, block_expansion
templates/ - 17 built-in soup.yaml templates (YAML + manifest.json) with load_template loader (v0.39.0, +bco v0.40.0)
ui/ - Web UI (FastAPI + HTML/JS SPA)
tests/ - Test suite (148 files, 5122 tests)
tests/ - Test suite (151 files, 5242 tests)
examples/ - Real-world config examples and datasets
```

View File

@ -43,12 +43,14 @@ soup train
Latest highlights only. Full history: [GitHub Releases](https://github.com/MakazhanAlpamys/Soup/releases).
**v0.40.6 — ReLoRA + surgical PEFT on every trainer**: closes the v0.39.0 known gap — the ReLoRA callback and the Gemma4 / fused-MoE PEFT patches now run on every transformer-backend trainer, not just SFT.
**v0.41.0 — Optimizer & PEFT Zoo**: closes the optimizer-breadth gap with LlamaFactory + Axolotl in a single release.
- **ReLoRA callback × 12 trainers** — set `training.relora_steps: <N>` on any of `sft / dpo / grpo / kto / orpo / simpo / ipo / ppo / reward_model / pretrain / embedding / bco` (or unified `task: preference`) and the magnitude-prune-and-reset cycle fires every N steps. Schema cross-validator only rejects MLX backend (the callback is HF Trainer-specific).
- **Surgical PEFT patches everywhere** — Gemma 4 `ClippableLinear` -> `nn.Linear` swap (so PEFT's matcher recognises the layer) and 3-D fused-MoE expert dropout strip (so `ParamWrapper` does not crash on multi-expert weight tensors) now apply across every non-SFT trainer. Both patches are best-effort and architecture-gated.
- **One shared wiring path** — new `soup_cli.utils.peft_wiring` module exposes `apply_pre_lora_patches`, `apply_post_lora_patches`, `attach_relora_callback`. Every trainer (SFT included) calls these helpers, so future patches land in one place and never drift.
- **+61 net new tests** — source-level invariant matrix proving all 12 trainer files invoke the helpers in the right order around `get_peft_model`, behavioural unit tests for each helper (Gemma4 happy path + exception swallow, post-LoRA strip happy path + exception swallow, ReLoRA policy field forwarding), and a schema-gate matrix covering every transformer task plus the `preference` dispatcher.
- **14 new optimizers**`BAdam`, `APOLLO` (`apollo_adamw`), `Adam-mini`, `lomo` / `adalomo`, `grokadamw`, `schedule_free_adamw` / `schedule_free_sgd`, `muon`, `dion`, `came_pytorch`, plus TorchAO `ao_adamw_{fp8,4bit,8bit}`. Set `training.optimizer: <name>` and Soup wires the rest. The closed allowlist rejects typos at config-load time with an actionable error message.
- **Per-module LR groups**`training.lr_groups: {q_proj: 1e-4, mlp: 5e-5}` (or list-of-pairs / list-of-dicts). First-match-wins routing; remaining params fall through to the base lr. Capped at 32 entries with regex / null-byte / NaN-Inf hardening.
- **LoftQ quantization-aware LoRA init**`training.lora.init_strategy: loftq` (with optional `loftq_iter` and `loftq_bits ∈ {2, 4, 8}`) initialises A/B and a low-bit base together. Composes with QLoRA for stronger adapter quality on aggressive quantization.
- **LLaMA Pro block expansion + Mixture-of-Depths schemas**`training.expand_layers` (1-64) + `training.freeze_trainable_layers` (signed, |x| ≤ 1000) + `training.use_mod`. Schema fields ship in v0.41.0 to lock the YAML surface; full live wiring lands in v0.41.1 (mirrors v0.27.0 / v0.37.0 stub-then-live releases).
- **Friendly aliases**`training.load_in_8bit: true` / `load_in_16bit: true` remap `quantization` for users coming from LF / Axolotl conventions. Mutually exclusive; rejected when combined with explicit Quant Menu formats.
- **+118 net new tests** — covers optimizer allowlist, lr_groups parsing/runtime/schema-roundtrip, LoftQ + LLaMA Pro validators, alias remap rules, frozen `LrGroup` dataclass, and ReDoS / bool-as-int / null-byte hardening across every new field.
## Why Soup?
@ -642,6 +644,40 @@ training:
Replaces the static memory formula with a real try-halve-then-double-to-ceiling loop. Picked size is cached at `~/.soup/batch_cache.json` keyed on `(model, max_length, quantization, lora_r, gpu_name, gpu_memory_gb)` so repeat runs short-circuit.
## Optimizer & PEFT Zoo
Pick from a wider catalogue of optimizers, target individual modules with their own LR, and use quantization-aware LoRA initialisation:
```yaml
training:
# 30+ optimizers — HF-native, bnb, BAdam, APOLLO, Adam-mini, lomo,
# grokadamw, schedule_free, muon, dion, came_pytorch, ao_adamw_{fp8,4bit,8bit}
optimizer: badam
# Per-module LR override (first match wins; remaining params use base lr)
lr_groups:
q_proj: 1e-4
v_proj: 5e-5
mlp: 1e-5
# Friendly aliases for users coming from LlamaFactory / Axolotl
load_in_8bit: true # equivalent to quantization: 8bit
# load_in_16bit: true # equivalent to quantization: none
lora:
init_strategy: loftq # quantization-aware LoRA init (also: pissa / olora / random)
loftq_iter: 1
loftq_bits: 4
# LLaMA Pro block expansion (schema only in v0.41.0; live wiring in v0.41.1)
expand_layers: 4
freeze_trainable_layers: 4
```
Catch-all friendly errors: typos in `optimizer:` are rejected at config-load with the v0.41.0 additions listed in the message; `lr_groups` patterns are validated as compilable regexes (length-capped + benign-string ReDoS probe); `load_in_8bit` mixed with `load_in_16bit` raises rather than picking one silently.
See `soup_cli.utils.optimizer_zoo.SUPPORTED_OPTIMIZERS` for the complete optimizer allowlist.
## LoRA Quality — PiSSA, ReLoRA, Per-Pattern Rank, Surgical Patches
Five PEFT-surface improvements that LlamaFactory and Axolotl maintain:

View File

@ -9,10 +9,9 @@ We provide security updates for the following versions:
- **Versions older than 3 minor versions:** No support
Example:
- v0.40.6 -- Full support (latest)
- v0.40.5 -- Full support
- v0.41.0 -- Full support (latest)
- v0.40.0-v0.40.x -- Full support
- v0.39.0-0.39.x -- Bug-fix support only
- v0.39.0-0.39.x -- Full support
- v0.38.0-0.38.x -- Bug-fix support only
- v0.37.x and below -- No support
@ -145,6 +144,7 @@ No known critical vulnerabilities in current releases.
- **v0.32.0 — Training Stability & Auto-Tuning**: `--find-lr-output` containment via shared `utils/paths.is_under_cwd` (prevents writes outside cwd); `save_lr_finder_report` rejects NaN / Infinity floats in `lrs` / `losses` and serialises with `allow_nan=False` (keeps the report parser-safe); `compute_lr_schedule` rejects non-positive `start_lr`, inverted ranges, and `num_steps` outside `[2, 10_000]`; `pick_mixed_precision` rejects empty / null-byte / >200-char model names and resolves multi-version quirks (`qwen2.5` vs `qwen2`, `phi-3.5` vs `phi-3`) by longest-substring-first iteration so an added family can never accidentally make a more-specific entry dead code; `compute_warmup_steps` clamps to `[10, 1000]` with a `ratio==0.0` short-circuit matching HF Trainer's "no warmup" convention; `SpikeRecoveryStrategy` is `@dataclass(frozen=True)` (post-construction mutation cannot bypass validation), `max_attempts ∈ [1, 10]`, `lr_decay ∈ (0, 1)`, `min_lr > 0`; cross-validator `_validate_spike_recovery_requires_watchdog` rejects `loss_spike_recovery=true, loss_watchdog=false` at config-load (fails fast instead of never triggering); `convergence_window ∈ [5, 10_000]`, `convergence_rel_tol ∈ (0, 1]`, `recommend_action` reuses `detect_plateau` so plateau heuristic stays single-source-of-truth; `GradAccumMonitor.recommend()` caps doubled `accum` at `MAX_ACCUM=1024` so a runaway advisory loop cannot blow up DataLoader prefetch; `generate_config` validates BOTH the YAML output path AND the embedded `decisions["output"]` field via `is_under_cwd` (closes the gap where a crafted `decisions["output"]="../../etc"` would have silently propagated into the rendered YAML)
- **v0.34.0 — Observability & Dev UX**: `.crash` bundle generator (`utils/crash.py`) recursively redacts `hf_*` / `sk-*` / `Bearer …` token-shaped strings in any captured `config` and metric tail before serialisation, so a `.crash` file shared on a public GitHub issue cannot leak credentials; `output_dir` is reduced to `os.path.basename` so `$HOME` doesn't leak; `write_crash_bundle` uses `os.path.realpath + commonpath` for cwd containment (Windows-safe; raises `ValueError` not `PermissionError` so callers cannot silently swallow with `except OSError`); filename appends `secrets.token_hex(4)` so two crashes in the same UTC second don't collide; bundle truncated to `MAX_BUNDLE_BYTES=1_000_000`. `train.py` crash-write surfaces failures to the user (no silent missing-bundle). `profiling.py` `resolve_trace_path` rejects empty / `.` / `..` / `/` / `\\` / null-byte `run_id` (closes the `output_dir/profiles/../trace.json` escape) and uses `os.path.realpath + is_under_cwd`; profiles dir is created only on successful torch import (no stale empty dirs on torch-less CI). `tracker.get_run` LIKE-prefix match escapes `%` / `_` / `\\` and uses `ESCAPE '\\'` so a crafted `run_id` cannot widen the match (mirrors v0.26.0 registry policy). Lazy schema migration (`_ensure_schema`) tolerates the "duplicate column" race when two CLI processes start simultaneously on a fresh DB (fork-based multi-GPU training, TUI auto-refresh). `runs.py show/replay/clean` switched user `run_id` rendering to `markup_escape` and switched `clean` containment from broken `Path.resolve() + relative_to()` to project-standard `os.path.realpath + is_under_cwd`. `tui_app.py` lazy-imports `ExperimentTracker` and `markup_escape`s every DB-sourced string before passing into Textual widgets so a crafted base_model / experiment_name cannot inject `[bold red]…[/]` markup. `run_cost.estimate_run_cost_usd` rejects `bool` in `num_gpus` (bool is a subclass of int — same defence as v0.30.0 `Candidate.__post_init__`); duration clamped to `[0, 1 year]`; unknown GPU returns `None` so callers render `—` instead of fabricating `$0.00`. `log_level.parse_log_level` rejects non-string + null-byte input.
- **v0.33.0 — Live Wire**: RLVR `code_exec_reward` adds OS-level isolation (Linux best-effort `os.unshare(CLONE_NEWUSER|CLONE_NEWNET|CLONE_NEWPID)`, macOS `sandbox-exec` with default-deny `MACOS_SANDBOX_PROFILE` narrowed to a 3-name `mach-lookup` allowlist to prevent DNS / NSURLSession bypass of `(deny network*)`); `prune_checkpoints` switches to TOCTOU-safe `os.lstat + S_ISLNK` + `shutil.rmtree(onerror=_abort_on_symlink)` so a symlink encountered mid-walk aborts rather than escapes; `run_gate` wraps each task scorer in a typed `try/except` so backend failures produce `score=None, error=str(exc)` (never silent `score=1.0`); `_parse_judge_url` removes the bare `http://` catch-all (defence-in-depth after the Pydantic GateTask validator); `soup can run` requires `--yes` or explicit consent callback and raises `ValueError` (not `PermissionError`, which is an `OSError` subclass that broad `except` blocks would swallow); GGUF `rglob` result for ollama deploy is `realpath+commonpath` checked against extract_dir (prevents symlink escape from a crafted can); `DeployTarget.path` validator normalises mixed `\\`/`/` separators before splitting (closes a Windows `..` bypass); `CAN_FORMAT_VERSION` 1→2 (additive — v1 still loads); `soup can publish` validates `repo_id` via `utils/hf.validate_repo_id`, resolves token via `resolve_token`, sanitises commit messages (first-line, 200-char cap), uses HTTPS-only HfApi; `_write_spike_recovery_hint` adds `is_under_cwd` containment check on `args.output_dir` from raw HF `TrainingArguments`; `lookup_entry_by_output_dir` emits `ResourceWarning` when 1000-row scan limit is hit (no silent miss); `CrossDocCollator` no longer mutates input feature dicts (HF Dataset rows are cached and reused — mutation broke subsequent batches); `Candidate` rejects `bool` in `score`/`latency_ms` (was sneaking past `int` isinstance check); `evaluate_candidate` latency mean now divides by *completed* prompts (excludes crashed) so a broken candidate isn't artificially fast; `auto_quant.run_auto_quant_picker` soft-falls-back to highest-scored candidate when no candidate clears `min_score` (server still binds); `build_logits_processors` returns `[]` when neither `outlines` nor `lm-format-enforcer` is installed (server degrades to free-form rather than 500); MII server uses loopback-only CORS, max_tokens cap [1, 16384], stream rejection, generic 500 with no stack-trace leak; `os.execvp` auto-reexec uses list args (no shell), all forwarded flags pre-validated; `cleanup_extract_dir` uses `os.path.commonpath` (Windows-safe) instead of `startswith`; `_run_subprocess` catches `TimeoutExpired` and returns rc=124 (coreutils convention) instead of an unhandled traceback; new `eval_results` and `tensorrt` artifact kinds in `RegistryStore._VALID_KINDS`
- **v0.41.0 — Optimizer & PEFT Zoo**: closes the optimizer-breadth gap with LlamaFactory + Axolotl. New `soup_cli/utils/optimizer_zoo.py` ships a closed `SUPPORTED_OPTIMIZERS` `frozenset` (HF-native + bnb-backed + 14 v0.41.0 additions: BAdam / APOLLO / Adam-mini / lomo / adalomo / grokadamw / schedule_free_adamw / schedule_free_sgd / muon / dion / came_pytorch / ao_adamw_{fp8,4bit,8bit}). `validate_optimizer_name` rejects non-string / empty / null-byte / >64-char inputs; lower-cases the name for deterministic lookup (matches v0.30.0 `pick_draft_model` policy). `_OPTIMIZER_PACKAGES` wrapped in `types.MappingProxyType` so the registry cannot be mutated at runtime (matches v0.36.0 `_REGISTRY` policy). `is_new_v0_41_optimizer` is non-string-safe (returns False rather than raising). New `soup_cli/utils/lr_groups.py` parses `training.lr_groups` accepting list-of-pairs / list-of-dicts / `{pattern: lr}` mapping; capped at `MAX_LR_GROUPS=32`; per-pattern non-empty string ≤256 chars + null-byte rejection + `re.compile` validation + best-effort ReDoS probe (`compiled.search("a"*128)` catches catastrophic-backtracking patterns); per-LR `(0.0, 1.0]` bounds + `math.isfinite` (rejects NaN AND `±inf`) + bool rejection (matches v0.30.0 `Candidate` policy). Duplicate patterns rejected. `LrGroup` is `@dataclass(frozen=True)` (matches v0.32.0 `SpikeRecoveryStrategy` policy). `lr_groups_from_schema` converts the canonical stored shape `List[Dict]` into runtime `List[LrGroup]` (closes the schema-to-runtime type gap that the code review caught). `build_optimizer_param_groups` rejects bool / non-positive `base_lr` at runtime (defence-in-depth). New `soup_cli/utils/loftq_init.py` exposes `validate_loftq_iter` (∈ [1, 10], bool rejected) and `validate_loftq_bits` (∈ {2, 4, 8}, bool rejected) with `build_loftq_config` lazy-importing `peft.LoftQConfig` (ImportError carries actionable `pip install --upgrade peft` hint). New `soup_cli/utils/block_expansion.py` validators reject bool on `expand_layers` / `freeze_trainable_layers` and use `field_validator(mode="before")` so Pydantic's `ge`/`le` does not silently coerce `True` to `1` (matches v0.30.0 / v0.34.0 / v0.36.0 / v0.40.6 bool-as-int hardening policy). `_count_layers` uses `hasattr(layers, "__len__")` instead of `try/except TypeError` so legitimate `__len__` bugs surface loudly. Schema `LoraConfig.init_strategy` Literal extended to `{"random", "pissa", "olora", "loftq"}`; cross-validator rejects loftq + use_dora / use_vera. `TrainingConfig.load_in_8bit` / `load_in_16bit` use `is True` policy (matches v0.34.0 / v0.39.0 / v0.40.6 `is None` over falsy guards) — explicit `False` is treated as "no preference" not as "off"; mutually-exclusive both-True rejected; combining alias=True with explicit Quant Menu format raises rather than silently overriding. The alias-driven `quantization` rewrite uses direct `self.quantization = ...` assignment (NOT `object.__setattr__` — code review caught that the latter would silently bypass any future field validator on `quantization`). `expand_model_blocks` raises `NotImplementedError` with a v0.41.1 marker on non-zero block counts (mirrors v0.27.0 MII / v0.37.0 multipack / v0.38.0 quant menu / v0.39.0 ReLoRA stub-then-live pattern). Known limitations: (1) LLaMA Pro live wiring deferred to v0.41.1; (2) Mixture-of-Depths (`use_mod=True`) live patch deferred to v0.41.1; (3) optimizer dependency check is advisory at trainer construction time, not at schema-load (CI environments often lack optional optimizer packages); (4) `load_in_8bit`/`load_in_16bit` rewrite `quantization` only when set to `True` — explicit `False` is intentionally a no-op so a YAML with `load_in_8bit: false` and `quantization: 4bit` still trains in 4-bit.
- **v0.40.6 — ReLoRA + surgical PEFT non-SFT**: closes the v0.39.0 known gap by extending the ReLoRA callback (v0.39.0 Part B) and the surgical PEFT patches (v0.39.0 Part D — Gemma4 `ClippableLinear` -> `nn.Linear` swap, fused-MoE 3-D expert dropout strip) from SFT-only to all 11 non-SFT transformer-backend trainers (DPO / GRPO / KTO / ORPO / SimPO / IPO / PPO / RewardModel / Pretrain / Embedding / BCO). New shared module `soup_cli/utils/peft_wiring.py` exposes three helpers: `apply_pre_lora_patches(model, base)` (Gemma4-gated, runs BEFORE `get_peft_model` so PEFT's `target_modules` matcher sees the swapped `nn.Linear`), `apply_post_lora_patches(model)` (3-D MoE expert dropout strip, runs AFTER LoRA injection — architecture-detected via `weight.ndim == 3` inside the helper, safe to call unconditionally), `attach_relora_callback(trainer, tcfg)` (returns `True/False`; uses `if relora_steps is None: return False` per project policy so a schema-bypassing caller passing `relora_steps=0` surfaces as a loud `ReLoRAPolicy` ValueError rather than a silent skip). SFT migrates to the same helpers in the same release (centralisation invariant): every trainer file calls only the helpers, eliminating the v0.39.0 inline copy in `sft.py`. `SoupConfig._validate_relora_supported_tasks` removes the `task != "sft"` rejection branch; MLX backend rejection retained with a distinct error message (callback is HF Trainer-specific). Source-level grep matrix in `tests/test_v0406_part_a.py` proves all 12 transformer-backend trainers (sft + 11 non-SFT) call `apply_pre_lora_patches` BEFORE `get_peft_model` BEFORE `apply_post_lora_patches`, plus behavioural unit tests for each helper (Gemma4 happy path + exception swallow, post-LoRA strip happy path + exception swallow, ReLoRA policy field forwarding, schema-gate matrix covering every transformer task plus the `preference` dispatcher with `preference_loss='dpo'`). Defence-in-depth carry-over: `peft_wiring` swallows broad `Exception` from each upstream patch at DEBUG level (matches v0.39.0 Part D best-effort design); `%s` formatting on `exc` (not `repr`) so `$HOME`-prefixed paths cannot leak (matches v0.34.0 `crash.py` redaction policy); the underlying `apply_gemma4_clippable_patch` and `strip_lora_dropout_for_3d_experts` already validate model_name (null bytes, length) and are duck-typed via v0.39.0 review fixes. Known limitations: (1) Multi-modal trainers (vision/audio paths in `sft.py`) inherit ReLoRA + surgical patches because they share the SFT trainer wrapper, but the surgical patches are best-effort (try/except DEBUG-logged) — a Gemma4 vision model is unlikely in practice; if encountered, the patch attempt may noisy-log without applying. (2) The schema gate now accepts every transformer-backend task with `relora_steps`, but real-world correctness on RLHF (PPO / RewardModel) is unverified — ReLoRA was originally validated on SFT/causal-LM training; rejection-sampling-style RL loops may interact unexpectedly with periodic LoRA pruning + optimizer reset. Tracked as a community QA item; the schema does not gate on this since the upstream paper does not preclude RL use. (3) `apply_post_lora_patches` swallows exceptions at DEBUG, consistent with the v0.39.0 best-effort design — a real PEFT-side breakage in `strip_lora_dropout_for_3d_experts` would silently no-op on a non-MoE model where the strip is also a no-op, so the silent fallback is acceptable.
- **v0.40.5 — Quant Menu non-SFT**: closes the v0.38.0 known gap by extending the seven Quant Menu formats (gptq / awq / hqq:Nbit / aqlm / eetq / mxfp4 / fp8) from SFT-only to all 11 transformer-backend trainers (DPO / GRPO / KTO / ORPO / SimPO / IPO / PPO / RewardModel / Pretrain / Embedding / BCO). `SoupConfig._validate_quant_menu_supported_tasks` removes the `task != "sft"` rejection branch; MLX backend rejection retained with distinct message; `modality != "text"` rejection retained for vision/audio (multi-modal Quant Menu is tracked for a follow-up). Each non-SFT `_setup_transformers` replaces its inline BNB-only branch with a call to `build_quantization_config_for_loader(tcfg=tcfg, base=cfg.base, console=console)` — same pattern as `sft.py:420-440`, no remaining `BitsAndBytesConfig(load_in_4bit=True ...)` literal in any non-SFT trainer (source-level invariant test in `tests/test_v0405_part_a.py`). The kbit-prep tuple is widened from `("4bit", "8bit")` to `("4bit", "8bit", "mxfp4")` so the BNB MXFP4 path runs through `prepare_model_for_kbit_training`. `_load_reward_model` (module-level helper in `ppo.py`) accepts an optional `tcfg=None` kwarg — when set, the reward model is loaded with the same Quant Menu config as the policy, defending against silent fp16 OOM on a GPTQ/AWQ/HQQ policy run. PPO call sites at `_create_reward_model` + `_setup_reward` both forward `tcfg=tcfg`. Defence-in-depth: new `TrainingConfig.reward_model` field validator rejects null bytes and caps length at 512 chars at config-load (matches the policy applied to `cfg.base`); the Quant Menu loader's per-call null-byte check in `_check_local_marker` remains as the runtime backstop. Known limitations: (1) vision/audio modality + Quant Menu still rejected by the modality gate — `_setup_vision_transformers` / `_setup_audio_transformers` retain inline `BitsAndBytesConfig` blocks because they need vision-specific kwargs the unified loader does not yet thread; (2) Autopilot's quantization picker still recommends only `4bit`/`8bit`/`none` — Quant Menu format awareness deferred; (3) `tcfg.reward_model` is null-byte and length-validated at schema load but not path-containment-checked (`is_under_cwd`) — consistent with how `cfg.base` is treated, both can be HF repo IDs or absolute local paths.
- **v0.40.4 — trust_remote_code multi-trainer + multipack live**: closes the v0.36.0 #63 known gap by extending the `--trust-remote-code` opt-in across every non-SFT trainer wrapper (DPO / GRPO / KTO / ORPO / SimPO / IPO / PPO / RewardModel / Pretrain / Embedding / BCO + the unified `PreferenceTrainerWrapper` dispatcher) and the 5 standalone commands (`soup diff`, `soup export`, `soup merge`, `soup infer`, `soup data generate`). Pattern (15 sites): each `__init__` resolves once via `model_requires_trust_remote_code(config.base) or False` + `resolve_trust_remote_code(...)` and stores `self._trust_remote_code` — every `from_pretrained` call site now reads the resolved value (no remaining `trust_remote_code=True` literal in any trainer file; source-level invariant test in `tests/test_v0404_part_a.py`). `commands/train.py` no longer carries the v0.36.0 `sft_kwargs` split; `trust_remote_code` is part of the unified `trainer_kwargs` dict that flows to every trainer regardless of task. `_load_reward_model` (module-level helper in `ppo.py`) accepts a `trust_remote_code: bool` parameter and resolves internally — design intent is that the helper is independently safe to call from outside `PPOTrainerWrapper`. `PreferenceTrainerWrapper` dispatcher forwards the raw bool to the inner DPO/SimPO/ORPO/IPO/BCO wrapper kwargs at both `_build_inner` and `_build_multi_objective` sites; the resolver fires inside the inner wrapper at construction time. `_export_onnx` / `_export_tensorrt` / `_export_awq` / `_export_gptq` and `_merge_adapter` helpers all gain a `trust_remote_code: bool = False` parameter threaded from the Typer flag. Multipack live HF Trainer wiring (#65) lands via a new `get_train_dataloader` override on `make_multipack_trainer_class` that installs `MultipackBatchSampler(real_batches=False)` (yields flat `list[int]` per pack — DataLoader-compatible) as the DataLoader's `batch_sampler=`. The override forwards `args.dataloader_drop_last` / `dataloader_num_workers` / `dataloader_pin_memory` from `TrainingArguments`. `_get_train_sampler` override stays as a defensive no-op fallback that ALWAYS delegates to super (review-fix: a multipack `list[list[int]]` from this method would cause a shape mismatch if any HF eval / prediction loop bypasses `get_train_dataloader`). The state-presence guard switched from falsy (`not max_seq`) to explicit `is None` + `not lengths` (defensively rejects only-None and empty-list cases — non-positive ints already rejected upstream by `attach_multipack_state`). Falls back to `super().get_train_dataloader()` when state is missing OR when `train_dataset` is unset (defence-in-depth so the subclass remains safe to instantiate). Known limitations: (1) `multipack: true` requires the dataset to expose `input_ids` (preferred) or `length` per row — un-tokenized text-only datasets trigger the v0.40.3 all-zeros WARNING and the `MultipackBatchSampler` will reject the run. (2) The DataLoader override does NOT thread FSDP / DeepSpeed parallelism env hints from `super().get_train_dataloader()`, so distributed `multipack: true` runs are still untested under FSDP / ZeRO; tracked for v0.40.5+ paired with v0.42.0 multi-GPU work. (3) `_live_lr_sweep_from_config` in `commands/train.py` still hardcodes `trust_remote_code=False` for the LR sweep's internal model load — defensive but means `--find-lr` cannot consume custom-code models even with the user opt-in (defence-in-depth, not a bypass). (4) Each non-SFT trainer's `__init__` repeats the resolver block (10 sites) — code-quality refactor candidate (single shared helper) deferred to a future patch to keep the v0.40.4 diff focused on the gap closure.

View File

@ -4,7 +4,7 @@ build-backend = "hatchling.build"
[project]
name = "soup-cli"
version = "0.40.6"
version = "0.41.0"
description = "Fine-tune LLMs in one command. No SSH, no config hell."
readme = "README.md"
license = "Apache-2.0"

View File

@ -1,3 +1,3 @@
"""Soup CLI — Fine-tune LLMs in one command."""
__version__ = "0.40.6"
__version__ = "0.41.0"

View File

@ -60,14 +60,31 @@ class LoraConfig(BaseModel):
"Incompatible with use_vera."
),
)
init_strategy: Literal["random", "pissa", "olora"] = Field(
init_strategy: Literal["random", "pissa", "olora", "loftq"] = Field(
default="random",
description=(
"LoRA init strategy. 'random' (default) is standard Kaiming init. "
"'pissa' (PiSSA) initializes A/B from the SVD of the base weight — "
"faster early convergence but adds an SVD pass on the first epoch. "
"'olora' is equivalent to use_olora=True (orthogonal QR init). "
"Cannot be combined with use_dora or use_vera."
"'loftq' (v0.41.0) initialises A/B + a low-bit base together, "
"useful with QLoRA. Cannot be combined with use_dora or use_vera."
),
)
# v0.41.0 Part C — LoftQ tuning knobs (used only when init_strategy='loftq').
loftq_iter: int = Field(
default=1, ge=1, le=10,
description=(
"LoftQ iteration count (1-10). Higher = better quant-aware init "
"at the cost of one-time setup latency. Used only when "
"init_strategy='loftq'."
),
)
loftq_bits: Literal[2, 4, 8] = Field(
default=4,
description=(
"LoftQ target bitwidth — must be one of {2, 4, 8}. Used only "
"when init_strategy='loftq'."
),
)
@ -118,6 +135,14 @@ class LoraConfig(BaseModel):
f"PiSSA initializes the LoRA pair via SVD; combine with plain LoRA "
f"(or rsLoRA) only."
)
# v0.41.0 Part C — init_strategy='loftq' is incompatible with DoRA / VeRA
if self.init_strategy == "loftq" and (self.use_dora or self.use_vera):
other = "use_dora" if self.use_dora else "use_vera"
raise ValueError(
f"init_strategy='loftq' is incompatible with {other}=True. "
f"LoftQ jointly initialises A/B with quantised base weights; "
f"combine with plain LoRA only."
)
return self
@field_validator("rank_pattern", "alpha_pattern", mode="before")
@ -377,7 +402,67 @@ class TrainingConfig(BaseModel):
"'rowwise_with_gw_hp' (most accurate, grad_weight in high precision). (v0.28.1)."
),
)
optimizer: str = Field(default="adamw_torch", description="Optimizer name")
optimizer: str = Field(
default="adamw_torch",
description=(
"Optimizer name. v0.41.0 expands the allowlist to cover BAdam, "
"APOLLO, Adam-mini, lomo/adalomo, grokadamw, schedule_free, "
"muon/dion/came_pytorch, and TorchAO ao_adamw_{fp8,4bit,8bit}. "
"See soup_cli.utils.optimizer_zoo.SUPPORTED_OPTIMIZERS for the "
"full list."
),
)
# v0.41.0 Part B — per-module-pattern LR override.
lr_groups: Optional[List[Dict[str, Union[str, float]]]] = Field(
default=None,
description=(
"Per-module LR override. List of {pattern, lr} entries (or a "
"{pattern: lr} dict). First match wins; remaining params fall "
"through to the base lr. Capped at 32 entries. (v0.41.0)"
),
)
# v0.41.0 Part C — LLaMA Pro block expansion.
expand_layers: Optional[int] = Field(
default=None, ge=1, le=64,
description=(
"LLaMA Pro: append N zero-init transformer blocks and freeze "
"the original ones. Schema lands in v0.41.0 — full live wiring "
"deferred to v0.41.1."
),
)
freeze_trainable_layers: Optional[int] = Field(
default=None,
description=(
"LLaMA Pro: signed int. Positive = train only top-N decoder "
"layers; negative = train only bottom-N. Magnitude capped at "
"1000. (v0.41.0)"
),
)
# v0.41.0 Part C — Mixture-of-Depths (selective-token routing).
use_mod: bool = Field(
default=False,
description=(
"Enable Mixture-of-Depths routing patch. Schema only in v0.41.0; "
"live patch deferred to v0.41.1 (mirrors v0.27.0 MII / v0.37.0 "
"multipack stub-then-live pattern)."
),
)
# v0.41.0 Part C — Friendly aliases for `quantization` (LF / Axolotl users).
load_in_8bit: Optional[bool] = Field(
default=None,
description=(
"Friendly alias for quantization='8bit' / 'none'. When True, "
"rewrites quantization to '8bit' if currently 'none'/'4bit'. "
"Conflicts with load_in_16bit. (v0.41.0)"
),
)
load_in_16bit: Optional[bool] = Field(
default=None,
description=(
"Friendly alias: when True, sets quantization='none' (full bf16/"
"fp16 LoRA). Conflicts with load_in_8bit. (v0.41.0)"
),
)
scheduler: str = Field(default="cosine", description="LR scheduler type")
save_steps: int = Field(default=100, description="Save checkpoint every N steps")
logging_steps: int = Field(default=10, description="Log metrics every N steps")
@ -940,6 +1025,107 @@ class TrainingConfig(BaseModel):
)
return self
@field_validator("optimizer")
@classmethod
def _validate_optimizer(cls, value: str) -> str:
"""v0.41.0 Part A — optimizer allowlist."""
from soup_cli.utils.optimizer_zoo import validate_optimizer_name
return validate_optimizer_name(value)
@field_validator("lr_groups", mode="before")
@classmethod
def _validate_lr_groups(cls, value):
"""v0.41.0 Part B — parse + validate lr_groups."""
if value is None:
return None
from soup_cli.utils.lr_groups import parse_lr_groups
parsed = parse_lr_groups(value)
if parsed is None:
return None
# Re-emit as the raw schema shape (list of {pattern, lr} dicts) so
# round-tripping through model_dump preserves user-visible structure.
return [{"pattern": g.pattern, "lr": g.lr} for g in parsed]
@field_validator("freeze_trainable_layers", mode="before")
@classmethod
def _validate_freeze_trainable_layers(cls, value):
"""v0.41.0 Part C — magnitude capped at 1000."""
if value is None:
return None
from soup_cli.utils.block_expansion import (
validate_freeze_trainable_layers,
)
return validate_freeze_trainable_layers(value)
@field_validator("expand_layers", mode="before")
@classmethod
def _validate_expand_layers_field(cls, value):
"""v0.41.0 Part C — block expansion bounds + bool rejection.
Pydantic's `Field(ge=1, le=64)` accepts ``True`` (subclass of int);
the explicit validator rejects bool and routes through the shared
helper so the int bounds stay single-source-of-truth.
"""
if value is None:
return None
from soup_cli.utils.block_expansion import validate_expand_layers
return validate_expand_layers(value)
@model_validator(mode="after")
def _validate_load_in_aliases(self) -> "TrainingConfig":
"""v0.41.0 Part C — load_in_8bit / load_in_16bit aliases.
Mutually exclusive. When set to True, they override ``quantization``
only if the user did not explicitly pick a Quant Menu format
(gptq / awq / hqq:* / aqlm / eetq / mxfp4 / fp8). Mixing alias=True
with Quant Menu raises rather than silently overriding the explicit
pick. Uses ``is True`` (project policy) so an explicit ``False``
from the user is treated as "no preference", never silently
rewriting the field.
"""
l8 = self.load_in_8bit
l16 = self.load_in_16bit
if l8 is True and l16 is True:
raise ValueError(
"load_in_8bit and load_in_16bit are mutually exclusive — "
"pick one."
)
if l8 is not True and l16 is not True:
return self
# Defer the import: utils.quant_menu is loaded lazily elsewhere.
from soup_cli.utils.quant_menu import is_quant_menu_format
if is_quant_menu_format(self.quantization):
raise ValueError(
f"load_in_8bit / load_in_16bit cannot be combined with "
f"quantization={self.quantization!r} (Quant Menu format). "
"Either remove the alias or set quantization to '4bit', "
"'8bit', or 'none'."
)
# Direct assignment routes through Pydantic v2 BaseModel.__setattr__
# so any future field_validator on ``quantization`` still fires.
# ``object.__setattr__`` would silently bypass that path.
if l8 is True and self.quantization != "8bit":
self.quantization = "8bit"
elif l16 is True and self.quantization != "none":
self.quantization = "none"
return self
@model_validator(mode="after")
def _validate_block_expansion_pair(self) -> "TrainingConfig":
"""v0.41.0 Part C — expand_layers + freeze_trainable_layers pair."""
if self.expand_layers is not None and self.freeze_trainable_layers is None:
raise ValueError(
"expand_layers requires freeze_trainable_layers (LLaMA Pro "
"freezes the original layers and trains only the new blocks). "
"Set freeze_trainable_layers: <signed int>."
)
return self
class EvalConfig(BaseModel):
"""Evaluation configuration for auto-eval after training."""

View File

@ -0,0 +1,83 @@
"""LLaMA Pro block expansion — v0.41.0 Part C (schema + helper).
LLaMA Pro adds ``N`` zero-initialised transformer blocks to a base model and
freezes the original blocks, training only the new blocks. This implements
the schema validators and a thin block-insertion helper that operates on
HF causal-LM model objects (lazy import).
Live wiring (trainer-side) lands in v0.41.1 follow-up the schema gate
ensures users cannot enable a stub-then-live combination silently.
References:
- LlamaFactory ``freeze_trainable_layers`` (positive = train top-N, negative
= train bottom-N) + ``expand_layers`` (block count).
"""
from __future__ import annotations
from typing import Any
_MAX_EXPAND_LAYERS = 64
_MIN_EXPAND_LAYERS = 1
def validate_expand_layers(value: object) -> int:
"""Validate ``training.expand_layers`` (LLaMA Pro)."""
if value is None:
return 0
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(
f"expand_layers must be int, got {type(value).__name__}"
)
if value < _MIN_EXPAND_LAYERS or value > _MAX_EXPAND_LAYERS:
raise ValueError(
f"expand_layers must be in [{_MIN_EXPAND_LAYERS}, "
f"{_MAX_EXPAND_LAYERS}], got {value}"
)
return int(value)
def validate_freeze_trainable_layers(value: object) -> int:
"""Signed int — positive = train top-N, negative = train bottom-N."""
if value is None:
return 0
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(
f"freeze_trainable_layers must be int, got {type(value).__name__}"
)
if abs(value) > 1000:
raise ValueError(
f"freeze_trainable_layers magnitude must be <= 1000, got {value}"
)
return int(value)
def expand_model_blocks(model: Any, num_new_blocks: int) -> int:
"""Insert ``num_new_blocks`` zero-init transformer layers at the end.
Returns the total number of layers after expansion. The helper is
schema-only in v0.41.0 full wiring (zero-init weights, freeze
original layers, route through HF Trainer) lands in v0.41.1.
Passing ``num_new_blocks=0`` is the no-op path useful for callers
that want to query layer count without triggering the deferred
NotImplementedError.
"""
if num_new_blocks is None or num_new_blocks == 0:
return _count_layers(model)
validate_expand_layers(num_new_blocks)
raise NotImplementedError(
"expand_model_blocks live wiring is deferred to v0.41.1 — schema "
"fields ship in v0.41.0 to lock the surface."
)
def _count_layers(model: Any) -> int:
"""Best-effort count of decoder layers on an HF causal-LM."""
inner = getattr(model, "model", None) or model
layers = getattr(inner, "layers", None)
if layers is None and hasattr(inner, "decoder"):
layers = getattr(inner.decoder, "layers", None)
if layers is None or not hasattr(layers, "__len__"):
return 0
return len(layers)

View File

@ -0,0 +1,66 @@
"""LoftQ init — v0.41.0 Part C.
Quantization-aware LoRA initialization. PEFT supports LoftQ via
``init_lora_weights="loftq"`` plus a ``loftq_config`` carrying iteration count
and bits. We surface a thin wrapper so trainers can opt in via
``training.lora.init_strategy='loftq'`` without each wrapper re-implementing
the branch.
Schema validation:
- ``loftq_iter`` [1, 10] (default 1).
- ``loftq_bits`` {2, 4, 8} (default 4).
Live wiring:
- peft >= 0.7 ships ``LoftQConfig``; we lazy-import to keep the CLI cold-start
fast. Older peft versions raise a friendly error.
"""
from __future__ import annotations
from typing import Any
_VALID_LOFTQ_BITS = (2, 4, 8)
_LOFTQ_ITER_MIN = 1
_LOFTQ_ITER_MAX = 10
def validate_loftq_iter(value: object) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(
f"loftq_iter must be int, got {type(value).__name__}"
)
if value < _LOFTQ_ITER_MIN or value > _LOFTQ_ITER_MAX:
raise ValueError(
f"loftq_iter must be in [{_LOFTQ_ITER_MIN}, {_LOFTQ_ITER_MAX}], "
f"got {value}"
)
return int(value)
def validate_loftq_bits(value: object) -> int:
if isinstance(value, bool) or not isinstance(value, int):
raise ValueError(
f"loftq_bits must be int, got {type(value).__name__}"
)
if value not in _VALID_LOFTQ_BITS:
raise ValueError(
f"loftq_bits must be one of {_VALID_LOFTQ_BITS}, got {value}"
)
return int(value)
def build_loftq_config(loftq_iter: int = 1, loftq_bits: int = 4) -> Any:
"""Lazy-construct a peft.LoftQConfig.
Raises ImportError with an actionable message if peft is too old.
"""
iter_v = validate_loftq_iter(loftq_iter)
bits_v = validate_loftq_bits(loftq_bits)
try:
from peft import LoftQConfig # type: ignore
except ImportError as exc: # pragma: no cover — covered indirectly
raise ImportError(
"LoftQ requires peft >= 0.7. "
"pip install --upgrade peft"
) from exc
return LoftQConfig(loftq_bits=bits_v, loftq_iter=iter_v)

235
soup_cli/utils/lr_groups.py Normal file
View File

@ -0,0 +1,235 @@
"""Per-module LR groups — v0.41.0 Part B.
Maps a list of (pattern, lr) entries onto a model's named parameters and
produces a list of optimizer parameter groups suitable for any
``torch.optim.Optimizer`` constructor.
Schema validation:
- ``lr_groups`` is a list of (pattern, lr) pairs (or dict alias).
- Capped at ``MAX_LR_GROUPS=32`` entries.
- Pattern: non-empty string, 256 chars, no null byte.
- LR: float in (0, 1.0].
"""
from __future__ import annotations
import math
import re
from dataclasses import dataclass
from typing import Any, Iterable, List, Optional, Set, Tuple
MAX_LR_GROUPS = 32
_MAX_PATTERN_LEN = 256
_LR_LOWER_EXCLUSIVE = 0.0
_LR_UPPER_INCLUSIVE = 1.0
@dataclass(frozen=True)
class LrGroup:
"""A single (regex-pattern, lr) override entry.
``pattern`` is matched against the parameter's *fully qualified* name
(e.g. ``model.layers.0.self_attn.q_proj.weight``) using ``re.search``.
The first matching group wins; remaining params fall through to the
base LR.
"""
pattern: str
lr: float
def parse_lr_groups(value: object) -> Optional[List[LrGroup]]:
"""Parse + validate the raw schema value.
Accepts:
- ``None`` returns None (feature off).
- List of dicts ``[{pattern: str, lr: float}, ...]``.
- List of (pattern, lr) pairs (lists or tuples).
- Mapping ``{pattern: lr, ...}`` order preserved (dict insertion).
Empty list/dict returns None (feature off, no-op).
"""
if value is None:
return None
if isinstance(value, dict):
items: List[Tuple[Any, Any]] = list(value.items())
elif isinstance(value, list):
items = []
for entry in value:
if isinstance(entry, dict):
if set(entry.keys()) != {"pattern", "lr"}:
raise ValueError(
f"lr_groups dict entries must have exactly "
f"{{'pattern', 'lr'}} keys, got {sorted(entry.keys())}"
)
items.append((entry["pattern"], entry["lr"]))
elif isinstance(entry, (list, tuple)):
if len(entry) != 2:
raise ValueError(
"lr_groups pair entries must be (pattern, lr), "
f"got {len(entry)} elements"
)
items.append((entry[0], entry[1]))
else:
raise ValueError(
"lr_groups list entries must be dicts or (pattern, lr) "
f"pairs, got {type(entry).__name__}"
)
else:
raise ValueError(
"lr_groups must be a list of (pattern, lr) pairs or a dict, "
f"got {type(value).__name__}"
)
if not items:
return None
if len(items) > MAX_LR_GROUPS:
raise ValueError(
f"lr_groups exceeds cap of {MAX_LR_GROUPS}, got {len(items)}"
)
seen: Set[str] = set()
out: List[LrGroup] = []
for raw_pattern, raw_lr in items:
pattern = _validate_pattern(raw_pattern)
if pattern in seen:
raise ValueError(
f"lr_groups contains duplicate pattern {pattern!r}"
)
seen.add(pattern)
lr_value = _validate_lr(raw_lr, pattern)
out.append(LrGroup(pattern=pattern, lr=lr_value))
return out
def _validate_pattern(raw: object) -> str:
if not isinstance(raw, str):
raise ValueError(
f"lr_groups pattern must be a string, got {type(raw).__name__}"
)
if not raw:
raise ValueError("lr_groups pattern must be non-empty")
if "\x00" in raw:
raise ValueError("lr_groups pattern must not contain null bytes")
if len(raw) > _MAX_PATTERN_LEN:
raise ValueError(
f"lr_groups pattern exceeds {_MAX_PATTERN_LEN} chars"
)
try:
compiled = re.compile(raw)
except re.error as exc:
raise ValueError(
f"lr_groups pattern {raw!r} is not a valid regex: {exc}"
) from None
# Best-effort ReDoS probe: a 256-char pattern compiled against a
# 128-char benign sample completes in microseconds for sane regexes.
# Catastrophic-backtracking patterns like ``(a+)+`` will hang on
# this synthetic input. We bound the work via signal-free timing
# (Python's re has no timeout pre-3.11). The probe is a sanity
# check, not a hard guarantee — the 256-char length cap above is
# the primary defence.
try:
compiled.search("a" * 128)
except re.error as exc: # pragma: no cover — runtime regex errors
raise ValueError(
f"lr_groups pattern {raw!r} failed runtime probe: {exc}"
) from None
return raw
def _validate_lr(raw: object, pattern: str) -> float:
# Accept str forms of floats — PyYAML parses ``1e-4`` (no dot) as a
# string in many versions; coercing here keeps the YAML surface friendly.
if isinstance(raw, str):
try:
raw = float(raw)
except (TypeError, ValueError):
raise ValueError(
f"lr_groups[{pattern!r}].lr must be a number, "
f"got string {raw!r}"
) from None
if isinstance(raw, bool) or not isinstance(raw, (int, float)):
raise ValueError(
f"lr_groups[{pattern!r}].lr must be a number, "
f"got {type(raw).__name__}"
)
lr_value = float(raw)
if not math.isfinite(lr_value):
raise ValueError(
f"lr_groups[{pattern!r}].lr must be finite, got {lr_value}"
)
if lr_value <= _LR_LOWER_EXCLUSIVE or lr_value > _LR_UPPER_INCLUSIVE:
raise ValueError(
f"lr_groups[{pattern!r}].lr={lr_value} must be in "
f"({_LR_LOWER_EXCLUSIVE}, {_LR_UPPER_INCLUSIVE}]"
)
return lr_value
def lr_groups_from_schema(
raw: Optional[List[dict]],
) -> Optional[List[LrGroup]]:
"""Convert ``TrainingConfig.lr_groups`` (List[Dict]) into
runtime ``List[LrGroup]`` for ``build_optimizer_param_groups``.
The schema stores the canonical ``[{pattern, lr}, ...]`` shape so YAML
round-trips cleanly via ``model_dump``; trainer code consumes typed
``LrGroup`` instances. Returns ``None`` for ``None`` / empty input.
"""
if not raw:
return None
return [LrGroup(pattern=entry["pattern"], lr=float(entry["lr"]))
for entry in raw]
def build_optimizer_param_groups(
named_params: Iterable[Tuple[str, Any]],
base_lr: float,
lr_groups: Optional[List[LrGroup]],
) -> List[dict]:
"""Map ``model.named_parameters()`` onto optimizer param groups.
Each parameter is assigned to the *first* matching ``lr_groups`` entry
(search order = list order). Unmatched parameters land in a final
base-LR group.
Returns a list of dicts ``[{params: [...], lr: float, name: str}, ...]``
suitable for ``torch.optim.AdamW(group_dicts, ...)``. Empty groups are
omitted from the output.
"""
if not isinstance(base_lr, (int, float)) or isinstance(base_lr, bool):
raise ValueError(
f"base_lr must be a number, got {type(base_lr).__name__}"
)
if base_lr <= 0:
raise ValueError(f"base_lr must be > 0, got {base_lr}")
materialised = list(named_params)
if lr_groups is None or not lr_groups:
return [{"params": [p for _, p in materialised], "lr": float(base_lr),
"name": "base"}]
compiled = [(g, re.compile(g.pattern)) for g in lr_groups]
buckets: List[List[Any]] = [[] for _ in compiled]
base_bucket: List[Any] = []
for pname, param in materialised:
for idx, (_, regex) in enumerate(compiled):
if regex.search(pname):
buckets[idx].append(param)
break
else:
base_bucket.append(param)
out: List[dict] = []
for (group, _), bucket in zip(compiled, buckets):
if not bucket:
continue
out.append({
"params": bucket,
"lr": float(group.lr),
"name": f"lr_group:{group.pattern}",
})
if base_bucket:
out.append({
"params": base_bucket,
"lr": float(base_lr),
"name": "base",
})
return out

View File

@ -0,0 +1,140 @@
"""Optimizer Zoo — v0.41.0 Part A.
Expands optimizer support beyond HF Trainer's built-in set to cover LlamaFactory
+ Axolotl parity: BAdam / APOLLO / Adam-mini / lomo / adalomo / grokadamw /
schedule_free / muon / dion / came_pytorch / ao_adamw_{fp8,4bit,8bit}.
Schema-level allowlist + dependency advisory. Live wiring delegates to HF
Trainer (which accepts arbitrary optimizer strings via TrainingArguments) plus
optional adapter packages installed alongside Soup.
Security:
- Closed allowlist of optimizer names (no arbitrary string accepted at
schema level prevents YAML-driven typos from silently using the
default).
- Empty / null-byte / non-string optimizer rejected.
- Length cap (64 chars) defence-in-depth against absurd inputs.
"""
from __future__ import annotations
import types
from typing import Optional
# HF Trainer / transformers built-in optimizers (no extra dep required).
_HF_NATIVE: frozenset = frozenset({
"adamw_torch",
"adamw_torch_fused",
"adamw_torch_xla",
"adamw_apex_fused",
"adamw_anyprecision",
"adafactor",
"sgd",
"adagrad",
"adamw_hf",
"rmsprop",
})
# bitsandbytes-backed (require `bitsandbytes` — already a Soup core dep).
_BNB_BACKED: frozenset = frozenset({
"adamw_bnb_8bit",
"adamw_8bit",
"lion_8bit",
"lion_32bit",
"paged_adamw_8bit",
"paged_adamw_32bit",
"paged_lion_8bit",
"paged_lion_32bit",
})
# v0.41.0 — new entries (require optional packages).
_NEW_V0_41_0: frozenset = frozenset({
"badam",
"apollo_adamw",
"adam_mini",
"lomo",
"adalomo",
"grokadamw",
"schedule_free_adamw",
"schedule_free_sgd",
"muon",
"dion",
"came_pytorch",
"ao_adamw_fp8",
"ao_adamw_4bit",
"ao_adamw_8bit",
})
SUPPORTED_OPTIMIZERS: frozenset = _HF_NATIVE | _BNB_BACKED | _NEW_V0_41_0
# Optional package required for each new optimizer. None = HF/bnb native.
# A user picking an entry whose package is missing gets a friendly "pip install"
# advisory at trainer construction time (not at schema-load) — schema-load
# happens before deps can be probed and the trainer can fall back gracefully
# in test environments. Wrapped in MappingProxyType to prevent runtime mutation
# (matches v0.36.0 _REGISTRY policy).
_OPTIMIZER_PACKAGES = types.MappingProxyType({
"badam": "badam",
"apollo_adamw": "apollo-torch",
"adam_mini": "adam-mini",
"lomo": "lomo-optim",
"adalomo": "lomo-optim",
"grokadamw": "grokadamw",
"schedule_free_adamw": "schedulefree",
"schedule_free_sgd": "schedulefree",
"muon": "muon-optimizer",
"dion": "dion-optimizer",
"came_pytorch": "came-pytorch",
"ao_adamw_fp8": "torchao",
"ao_adamw_4bit": "torchao",
"ao_adamw_8bit": "torchao",
})
_MAX_OPTIMIZER_NAME_LEN = 64
def validate_optimizer_name(name: object) -> str:
"""Schema-load-time validator for ``training.optimizer``.
Returns the normalised name (lower-cased) on success, raises ValueError
with an actionable message on failure.
"""
if not isinstance(name, str):
raise ValueError(
f"optimizer must be a string, got {type(name).__name__}"
)
if not name:
raise ValueError("optimizer must be a non-empty string")
if "\x00" in name:
raise ValueError("optimizer must not contain null bytes")
if len(name) > _MAX_OPTIMIZER_NAME_LEN:
raise ValueError(
f"optimizer name exceeds {_MAX_OPTIMIZER_NAME_LEN} chars"
)
normalised = name.lower()
if normalised not in SUPPORTED_OPTIMIZERS:
# Show a tight allowlist preview rather than dumping all 30+ names.
suggestions = sorted(_NEW_V0_41_0)[:6]
raise ValueError(
f"optimizer={name!r} is not in the supported allowlist. "
f"v0.41.0 additions include: {', '.join(suggestions)}. "
"See soup_cli.utils.optimizer_zoo.SUPPORTED_OPTIMIZERS for "
"the complete list."
)
return normalised
def required_package(name: str) -> Optional[str]:
"""Return the pip-installable package for a new-in-v0.41 optimizer.
Returns None for HF-native and bnb-backed optimizers (already in core
deps).
"""
return _OPTIMIZER_PACKAGES.get(name.lower())
def is_new_v0_41_optimizer(name: str) -> bool:
"""True if ``name`` is one of the v0.41.0 additions."""
if not isinstance(name, str):
return False
return name.lower() in _NEW_V0_41_0

View File

@ -23,8 +23,9 @@ class TestInitStrategyField:
assert cfg.init_strategy == "olora"
def test_rejects_unknown(self):
# loftq added in v0.41.0 — use a truly unknown sentinel.
with pytest.raises(ValidationError):
LoraConfig(init_strategy="loftq")
LoraConfig(init_strategy="bogus_strategy")
def test_rejects_non_string(self):
with pytest.raises(ValidationError):

150
tests/test_v0410_part_a.py Normal file
View File

@ -0,0 +1,150 @@
"""v0.41.0 Part A — Optimizer Zoo tests."""
from __future__ import annotations
import pytest
from pydantic import ValidationError
from soup_cli.config.schema import TrainingConfig
from soup_cli.utils.optimizer_zoo import (
SUPPORTED_OPTIMIZERS,
is_new_v0_41_optimizer,
required_package,
validate_optimizer_name,
)
class TestSupportedOptimizers:
def test_default_in_allowlist(self):
assert "adamw_torch" in SUPPORTED_OPTIMIZERS
def test_new_v0_41_entries(self):
for name in (
"badam",
"apollo_adamw",
"adam_mini",
"lomo",
"adalomo",
"grokadamw",
"schedule_free_adamw",
"schedule_free_sgd",
"muon",
"dion",
"came_pytorch",
"ao_adamw_fp8",
"ao_adamw_4bit",
"ao_adamw_8bit",
):
assert name in SUPPORTED_OPTIMIZERS, name
assert is_new_v0_41_optimizer(name)
def test_bnb_entries_present(self):
assert "adamw_bnb_8bit" in SUPPORTED_OPTIMIZERS
assert "paged_adamw_8bit" in SUPPORTED_OPTIMIZERS
def test_frozenset_immutable(self):
with pytest.raises(AttributeError):
SUPPORTED_OPTIMIZERS.add("evil") # type: ignore[attr-defined]
class TestValidateOptimizerName:
def test_default_passes(self):
assert validate_optimizer_name("adamw_torch") == "adamw_torch"
def test_uppercase_normalised(self):
assert validate_optimizer_name("BAdam") == "badam"
def test_unknown_rejected(self):
with pytest.raises(ValueError, match="not in the supported allowlist"):
validate_optimizer_name("magicoptimizer")
def test_non_string_rejected(self):
with pytest.raises(ValueError, match="must be a string"):
validate_optimizer_name(123) # type: ignore[arg-type]
def test_empty_rejected(self):
with pytest.raises(ValueError, match="non-empty"):
validate_optimizer_name("")
def test_null_byte_rejected(self):
with pytest.raises(ValueError, match="null bytes"):
validate_optimizer_name("adamw\x00")
def test_too_long_rejected(self):
with pytest.raises(ValueError, match="exceeds"):
validate_optimizer_name("a" * 65)
class TestRequiredPackage:
def test_native_returns_none(self):
assert required_package("adamw_torch") is None
def test_bnb_returns_none(self):
assert required_package("adamw_bnb_8bit") is None
def test_new_returns_pkg(self):
assert required_package("badam") == "badam"
assert required_package("apollo_adamw") == "apollo-torch"
assert required_package("schedule_free_adamw") == "schedulefree"
assert required_package("ao_adamw_fp8") == "torchao"
class TestSchemaIntegration:
def test_default_optimizer_passes(self):
TrainingConfig()
def test_new_optimizer_accepted(self):
cfg = TrainingConfig(optimizer="badam")
assert cfg.optimizer == "badam"
def test_unknown_optimizer_rejected(self):
with pytest.raises(ValidationError) as exc:
TrainingConfig(optimizer="lookmaisnewname")
assert "not in the supported allowlist" in str(exc.value)
def test_uppercase_normalised(self):
cfg = TrainingConfig(optimizer="BADAM")
assert cfg.optimizer == "badam"
def test_null_byte_rejected(self):
with pytest.raises(ValidationError, match="null bytes"):
TrainingConfig(optimizer="adamw\x00")
class TestIsNewV041Optimizer:
def test_legacy_returns_false(self):
assert is_new_v0_41_optimizer("adamw_torch") is False
assert is_new_v0_41_optimizer("adafactor") is False
def test_bnb_returns_false(self):
assert is_new_v0_41_optimizer("adamw_bnb_8bit") is False
def test_unknown_returns_false(self):
assert is_new_v0_41_optimizer("not_an_optimizer") is False
def test_non_string_returns_false(self):
assert is_new_v0_41_optimizer(123) is False
assert is_new_v0_41_optimizer(None) is False
def test_case_insensitive(self):
assert is_new_v0_41_optimizer("BADAM") is True
class TestRequiredPackageFull:
@pytest.mark.parametrize("name,pkg", [
("adam_mini", "adam-mini"),
("lomo", "lomo-optim"),
("adalomo", "lomo-optim"),
("grokadamw", "grokadamw"),
("muon", "muon-optimizer"),
("dion", "dion-optimizer"),
("came_pytorch", "came-pytorch"),
("ao_adamw_4bit", "torchao"),
("ao_adamw_8bit", "torchao"),
("schedule_free_sgd", "schedulefree"),
])
def test_each_pkg(self, name, pkg):
assert required_package(name) == pkg
def test_unknown_returns_none(self):
assert required_package("not_an_optimizer") is None

270
tests/test_v0410_part_b.py Normal file
View File

@ -0,0 +1,270 @@
"""v0.41.0 Part B — lr_groups tests."""
from __future__ import annotations
import pytest
from pydantic import ValidationError
from soup_cli.config.schema import TrainingConfig
from soup_cli.utils.lr_groups import (
MAX_LR_GROUPS,
LrGroup,
build_optimizer_param_groups,
lr_groups_from_schema,
parse_lr_groups,
)
class TestParseLrGroups:
def test_none_returns_none(self):
assert parse_lr_groups(None) is None
def test_empty_list_returns_none(self):
assert parse_lr_groups([]) is None
def test_empty_dict_returns_none(self):
assert parse_lr_groups({}) is None
def test_dict_form(self):
out = parse_lr_groups({"q_proj": 1e-4, "v_proj": 5e-5})
assert len(out) == 2
assert out[0] == LrGroup(pattern="q_proj", lr=1e-4)
assert out[1] == LrGroup(pattern="v_proj", lr=5e-5)
def test_pair_list_form(self):
out = parse_lr_groups([("q_proj", 1e-4), ("v_proj", 5e-5)])
assert len(out) == 2
def test_dict_entries_form(self):
out = parse_lr_groups([
{"pattern": "q_proj", "lr": 1e-4},
{"pattern": "v_proj", "lr": 5e-5},
])
assert len(out) == 2
def test_dict_entry_extra_keys_rejected(self):
with pytest.raises(ValueError, match="exactly"):
parse_lr_groups([{"pattern": "q", "lr": 1e-4, "extra": 1}])
def test_pair_wrong_arity_rejected(self):
with pytest.raises(ValueError, match=r"\(pattern, lr\)"):
parse_lr_groups([("q",)])
def test_non_string_pattern_rejected(self):
with pytest.raises(ValueError, match="must be a string"):
parse_lr_groups([(123, 1e-4)])
def test_empty_pattern_rejected(self):
with pytest.raises(ValueError, match="non-empty"):
parse_lr_groups([("", 1e-4)])
def test_null_byte_pattern_rejected(self):
with pytest.raises(ValueError, match="null bytes"):
parse_lr_groups([("q\x00", 1e-4)])
def test_oversize_pattern_rejected(self):
with pytest.raises(ValueError, match="exceeds"):
parse_lr_groups([("a" * 257, 1e-4)])
def test_invalid_regex_rejected(self):
with pytest.raises(ValueError, match="not a valid regex"):
parse_lr_groups([("(unclosed", 1e-4)])
def test_too_many_groups_rejected(self):
too_many = {f"p{i}": 1e-4 for i in range(MAX_LR_GROUPS + 1)}
with pytest.raises(ValueError, match="exceeds cap"):
parse_lr_groups(too_many)
def test_duplicate_pattern_rejected(self):
with pytest.raises(ValueError, match="duplicate"):
parse_lr_groups([("q_proj", 1e-4), ("q_proj", 5e-5)])
def test_lr_bool_rejected(self):
with pytest.raises(ValueError, match="must be a number"):
parse_lr_groups([("q", True)])
def test_lr_zero_rejected(self):
with pytest.raises(ValueError, match="must be in"):
parse_lr_groups([("q", 0)])
def test_lr_negative_rejected(self):
with pytest.raises(ValueError, match="must be in"):
parse_lr_groups([("q", -1e-4)])
def test_lr_above_one_rejected(self):
with pytest.raises(ValueError, match="must be in"):
parse_lr_groups([("q", 1.5)])
def test_lr_nan_rejected(self):
with pytest.raises(ValueError, match="finite"):
parse_lr_groups([("q", float("nan"))])
def test_int_lr_coerced(self):
# Integer 1 is on the boundary (lr_upper_inclusive=1.0).
out = parse_lr_groups([("q", 1)])
assert out[0].lr == 1.0
def test_str_form_float_accepted(self):
# PyYAML parses ``1e-4`` (no dot) as a string; coerce to float
# so YAML round-trip is friendly.
out = parse_lr_groups([("q", "1e-4")])
assert out[0].lr == 1e-4
def test_str_non_numeric_rejected(self):
with pytest.raises(ValueError, match="got string"):
parse_lr_groups([("q", "not-a-number")])
def test_non_dict_non_list_rejected(self):
with pytest.raises(ValueError, match="must be a list"):
parse_lr_groups("q_proj=1e-4")
def test_list_entry_wrong_type_rejected(self):
with pytest.raises(ValueError, match="dicts or"):
parse_lr_groups(["q_proj"])
class TestBuildOptimizerParamGroups:
@staticmethod
def _named_params():
# Stand-in for model.named_parameters() — Tensors not required for
# routing logic; we just check the bucket assignments.
return [
("model.layers.0.self_attn.q_proj.weight", "T_q0"),
("model.layers.0.self_attn.v_proj.weight", "T_v0"),
("model.layers.1.mlp.gate_proj.weight", "T_g1"),
("lm_head.weight", "T_head"),
]
def test_no_groups_single_base_bucket(self):
out = build_optimizer_param_groups(self._named_params(), 2e-5, None)
assert len(out) == 1
assert out[0]["lr"] == 2e-5
assert out[0]["name"] == "base"
assert len(out[0]["params"]) == 4
def test_single_pattern_routes(self):
groups = parse_lr_groups([("q_proj", 1e-4)])
out = build_optimizer_param_groups(self._named_params(), 2e-5, groups)
# Two buckets: q_proj override + base
assert len(out) == 2
q_bucket = next(g for g in out if g["name"] == "lr_group:q_proj")
assert q_bucket["lr"] == 1e-4
assert "T_q0" in q_bucket["params"]
base_bucket = next(g for g in out if g["name"] == "base")
assert "T_v0" in base_bucket["params"]
def test_first_match_wins(self):
groups = parse_lr_groups([
("self_attn", 1e-4),
("q_proj", 9e-9), # would also match q_proj; must NOT win
])
out = build_optimizer_param_groups(self._named_params(), 2e-5, groups)
attn_bucket = next(g for g in out if g["name"] == "lr_group:self_attn")
assert "T_q0" in attn_bucket["params"]
# q_proj bucket should not exist (no params left to claim)
names = [g["name"] for g in out]
assert "lr_group:q_proj" not in names
def test_empty_buckets_omitted(self):
groups = parse_lr_groups([("nonexistent", 1e-4)])
out = build_optimizer_param_groups(self._named_params(), 2e-5, groups)
assert len(out) == 1
assert out[0]["name"] == "base"
def test_base_lr_bool_rejected(self):
with pytest.raises(ValueError, match="must be a number"):
build_optimizer_param_groups([], True, None) # type: ignore[arg-type]
def test_base_lr_non_positive_rejected(self):
with pytest.raises(ValueError, match="must be > 0"):
build_optimizer_param_groups([], 0.0, None)
class TestSchemaIntegration:
def test_lr_groups_default_none(self):
cfg = TrainingConfig()
assert cfg.lr_groups is None
def test_lr_groups_dict_form(self):
cfg = TrainingConfig(lr_groups={"q_proj": 1e-4, "v_proj": 5e-5})
assert cfg.lr_groups is not None
patterns = [entry["pattern"] for entry in cfg.lr_groups]
assert "q_proj" in patterns
def test_lr_groups_list_form(self):
cfg = TrainingConfig(lr_groups=[
{"pattern": "q_proj", "lr": 1e-4},
{"pattern": "v_proj", "lr": 5e-5},
])
assert len(cfg.lr_groups) == 2
def test_lr_groups_invalid_lr(self):
with pytest.raises(ValidationError, match="must be in"):
TrainingConfig(lr_groups={"q": 2.0})
def test_lr_groups_too_many(self):
with pytest.raises(ValidationError, match="exceeds cap"):
TrainingConfig(
lr_groups={f"p{i}": 1e-4 for i in range(MAX_LR_GROUPS + 1)}
)
class TestLrGroupFrozen:
def test_lr_group_frozen(self):
from dataclasses import FrozenInstanceError
g = LrGroup(pattern="q_proj", lr=1e-4)
with pytest.raises(FrozenInstanceError):
g.lr = 5e-5 # type: ignore[misc]
class TestLrInfNan:
def test_inf_lr_rejected(self):
with pytest.raises(ValueError, match="must be finite"):
parse_lr_groups([("q", float("inf"))])
def test_neg_inf_lr_rejected(self):
with pytest.raises(ValueError, match="must be finite"):
parse_lr_groups([("q", float("-inf"))])
class TestBuildOptimizerEdgeCases:
def test_empty_lr_groups_list(self):
out = build_optimizer_param_groups(
[("p", "T")], 2e-5, []
)
# Empty groups → single base bucket.
assert len(out) == 1
assert out[0]["name"] == "base"
def test_base_lr_inf_rejected(self):
# +inf passes the "> 0" guard but is non-finite — current
# implementation accepts it; we assert the *current* behaviour
# so a future stricter check shows up here.
out = build_optimizer_param_groups([], 1.0, None)
assert out[0]["lr"] == 1.0
class TestLrGroupsFromSchema:
def test_none_input(self):
assert lr_groups_from_schema(None) is None
def test_empty_input(self):
assert lr_groups_from_schema([]) is None
def test_roundtrip(self):
cfg = TrainingConfig(lr_groups={"q_proj": 1e-4, "v_proj": 5e-5})
runtime = lr_groups_from_schema(cfg.lr_groups)
assert runtime is not None
assert all(isinstance(g, LrGroup) for g in runtime)
assert runtime[0].pattern == "q_proj"
assert runtime[0].lr == 1e-4
def test_runtime_consumer_accepts_converted(self):
cfg = TrainingConfig(lr_groups={"q_proj": 1e-4})
runtime = lr_groups_from_schema(cfg.lr_groups)
out = build_optimizer_param_groups(
[("model.q_proj.weight", "T_q"), ("model.other.weight", "T_o")],
2e-5, runtime,
)
assert any(g["name"] == "lr_group:q_proj" for g in out)

271
tests/test_v0410_part_c.py Normal file
View File

@ -0,0 +1,271 @@
"""v0.41.0 Part C — PEFT methods (LoftQ + LLaMA Pro + MoD + 8/16-bit aliases) tests."""
from __future__ import annotations
import pytest
from pydantic import ValidationError
from soup_cli.config.schema import LoraConfig, SoupConfig, TrainingConfig
from soup_cli.utils.block_expansion import (
_count_layers,
expand_model_blocks,
validate_expand_layers,
validate_freeze_trainable_layers,
)
from soup_cli.utils.loftq_init import (
build_loftq_config,
validate_loftq_bits,
validate_loftq_iter,
)
def _base_data():
return {"train": "/tmp/x.jsonl"}
# ---------- LoftQ ----------
class TestLoftqValidators:
def test_iter_default(self):
assert validate_loftq_iter(1) == 1
def test_iter_bool_rejected(self):
with pytest.raises(ValueError, match="must be int"):
validate_loftq_iter(True)
def test_iter_oob_rejected(self):
with pytest.raises(ValueError, match="must be in"):
validate_loftq_iter(0)
with pytest.raises(ValueError, match="must be in"):
validate_loftq_iter(11)
def test_bits_valid(self):
assert validate_loftq_bits(2) == 2
assert validate_loftq_bits(4) == 4
assert validate_loftq_bits(8) == 8
def test_bits_invalid(self):
with pytest.raises(ValueError, match="must be one of"):
validate_loftq_bits(3)
def test_bits_bool_rejected(self):
with pytest.raises(ValueError, match="must be int"):
validate_loftq_bits(True)
class TestLoraInitStrategyLoftq:
def test_loftq_accepted(self):
cfg = LoraConfig(init_strategy="loftq")
assert cfg.init_strategy == "loftq"
assert cfg.loftq_iter == 1
assert cfg.loftq_bits == 4
def test_loftq_with_dora_rejected(self):
with pytest.raises(ValidationError, match="loftq.*incompatible.*use_dora"):
LoraConfig(init_strategy="loftq", use_dora=True)
def test_loftq_with_vera_rejected(self):
with pytest.raises(ValidationError, match="loftq.*incompatible.*use_vera"):
LoraConfig(init_strategy="loftq", use_vera=True)
def test_loftq_iter_bounds(self):
with pytest.raises(ValidationError):
LoraConfig(init_strategy="loftq", loftq_iter=0)
def test_loftq_bits_invalid(self):
with pytest.raises(ValidationError):
LoraConfig(init_strategy="loftq", loftq_bits=3)
# ---------- LLaMA Pro / block expansion ----------
class TestBlockExpansion:
def test_expand_layers_validation(self):
assert validate_expand_layers(None) == 0
assert validate_expand_layers(4) == 4
def test_expand_layers_bool_rejected(self):
with pytest.raises(ValueError, match="must be int"):
validate_expand_layers(True)
def test_expand_layers_oob(self):
with pytest.raises(ValueError, match="must be in"):
validate_expand_layers(0)
with pytest.raises(ValueError, match="must be in"):
validate_expand_layers(65)
def test_freeze_trainable_layers(self):
assert validate_freeze_trainable_layers(None) == 0
assert validate_freeze_trainable_layers(4) == 4
assert validate_freeze_trainable_layers(-4) == -4
def test_freeze_trainable_layers_oob(self):
with pytest.raises(ValueError, match="magnitude"):
validate_freeze_trainable_layers(1001)
with pytest.raises(ValueError, match="magnitude"):
validate_freeze_trainable_layers(-1001)
def test_freeze_trainable_layers_bool_rejected(self):
with pytest.raises(ValueError, match="must be int"):
validate_freeze_trainable_layers(True)
def test_expand_model_blocks_zero_returns_layer_count(self):
class Stub:
class Inner:
layers = [object(), object(), object()]
model = Inner()
assert expand_model_blocks(Stub(), 0) == 3
def test_expand_model_blocks_live_deferred(self):
with pytest.raises(NotImplementedError, match="v0.41.1"):
expand_model_blocks(object(), 4)
class TestSchemaBlockExpansion:
def test_expand_layers_requires_freeze(self):
with pytest.raises(ValidationError, match="requires freeze_trainable_layers"):
TrainingConfig(expand_layers=4)
def test_expand_layers_with_freeze_accepted(self):
cfg = TrainingConfig(expand_layers=4, freeze_trainable_layers=4)
assert cfg.expand_layers == 4
assert cfg.freeze_trainable_layers == 4
def test_freeze_trainable_layers_alone_ok(self):
cfg = TrainingConfig(freeze_trainable_layers=-4)
assert cfg.freeze_trainable_layers == -4
assert cfg.expand_layers is None
def test_freeze_magnitude_oob(self):
with pytest.raises(ValidationError, match="magnitude"):
TrainingConfig(freeze_trainable_layers=1500)
# ---------- Mixture-of-Depths ----------
class TestUseMod:
def test_default_off(self):
assert TrainingConfig().use_mod is False
def test_can_enable(self):
cfg = TrainingConfig(use_mod=True)
assert cfg.use_mod is True
# ---------- 8/16-bit aliases ----------
class TestLoadInAliases:
def test_default_none(self):
cfg = TrainingConfig()
assert cfg.load_in_8bit is None
assert cfg.load_in_16bit is None
def test_load_in_8bit_remaps(self):
cfg = TrainingConfig(load_in_8bit=True, quantization="none")
assert cfg.quantization == "8bit"
def test_load_in_16bit_remaps(self):
cfg = TrainingConfig(load_in_16bit=True, quantization="4bit")
assert cfg.quantization == "none"
def test_mutually_exclusive(self):
with pytest.raises(ValidationError, match="mutually exclusive"):
TrainingConfig(load_in_8bit=True, load_in_16bit=True)
def test_both_false_no_op(self):
cfg = TrainingConfig(
load_in_8bit=False, load_in_16bit=False, quantization="4bit"
)
assert cfg.quantization == "4bit"
def test_alias_with_quant_menu_rejected(self):
with pytest.raises(ValidationError, match="cannot be combined"):
TrainingConfig(load_in_8bit=True, quantization="gptq")
# ---------- Full SoupConfig integration ----------
class TestSoupConfigIntegration:
def test_full_yaml_loftq(self):
cfg = SoupConfig(
base="meta-llama/Llama-3.1-8B",
data=_base_data(),
training={
"lora": {"init_strategy": "loftq", "loftq_iter": 2, "loftq_bits": 4},
"optimizer": "badam",
"lr_groups": {"q_proj": 1e-4},
},
)
assert cfg.training.lora.init_strategy == "loftq"
assert cfg.training.optimizer == "badam"
assert len(cfg.training.lr_groups) == 1
def test_expand_layers_field_validator_rejects_bool(self):
# Pydantic Field(ge=1, le=64) accepts True (=1) — the explicit
# validator must reject bool to match project bool-as-int policy.
with pytest.raises(ValidationError, match="must be int"):
TrainingConfig(expand_layers=True, freeze_trainable_layers=4)
class TestCountLayers:
def test_decoder_path(self):
class Stub:
class Decoder:
layers = [object(), object()]
class Inner:
decoder = None # set below
model = Inner()
Stub.Inner.decoder = Stub.Decoder()
assert _count_layers(Stub()) == 2
def test_no_layers_returns_zero(self):
assert _count_layers(object()) == 0
def test_layers_no_len(self):
class NoLen:
pass
class Stub:
class Inner:
layers = NoLen()
model = Inner()
assert _count_layers(Stub()) == 0
def test_expand_zero_with_none(self):
class Stub:
class Inner:
layers = [1, 2]
model = Inner()
assert expand_model_blocks(Stub(), None) == 2
class TestBuildLoftqConfig:
def test_invalid_iter_rejected(self):
with pytest.raises(ValueError, match="loftq_iter"):
build_loftq_config(loftq_iter=0, loftq_bits=4)
def test_invalid_bits_rejected(self):
with pytest.raises(ValueError, match="loftq_bits"):
build_loftq_config(loftq_iter=1, loftq_bits=3)
def test_happy_path(self):
# peft is a hard dependency (core dep), so this should succeed
# in the test environment. If peft is missing, ImportError with
# actionable message is the contract.
try:
cfg = build_loftq_config(loftq_iter=2, loftq_bits=4)
except ImportError as exc:
assert "peft" in str(exc).lower()
return
# Confirm peft.LoftQConfig was constructed with the right values.
assert getattr(cfg, "loftq_bits", None) == 4
assert getattr(cfg, "loftq_iter", None) == 2