mirror of https://github.com/razor-ai/soup.git
Phase 2.5: add export GGUF, merge LoRA, resume training, W&B integration (v0.2.0)
New commands: - soup export --model ./output --format gguf --quant q4_k_m - soup merge --adapter ./output New train flags: - soup train --resume auto (or --resume ./checkpoint-500) - soup train --wandb 184 tests passing (was 147), all lint clean. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
parent
2f3cf90d7b
commit
f03b578428
13
CLAUDE.md
13
CLAUDE.md
|
|
@ -53,6 +53,14 @@ soup train --config soup.yaml
|
|||
|
||||
**Eval:** `commands/eval.py` wraps lm-evaluation-harness for model evaluation on standard benchmarks (mmlu, gsm8k, etc.) with results saved to the experiment tracker.
|
||||
|
||||
**Merge:** `commands/merge.py` merges a LoRA adapter with its base model into a full standalone model using peft's `merge_and_unload()`. Auto-detects base model from `adapter_config.json`.
|
||||
|
||||
**Export:** `commands/export.py` exports models to GGUF format for Ollama/llama.cpp. Handles LoRA adapters (auto-merge first), then uses llama.cpp's `convert_hf_to_gguf.py` script for conversion and optional quantization. Auto-clones llama.cpp to `~/.soup/llama.cpp` if needed.
|
||||
|
||||
**Resume training:** `commands/train.py` supports `--resume auto` (find latest checkpoint) or `--resume <path>` to continue training from a checkpoint. Passes `resume_from_checkpoint` to HF Trainer.
|
||||
|
||||
**W&B integration:** `commands/train.py` supports `--wandb` flag to enable Weights & Biases logging. Sets `report_to="wandb"` in TrainingArguments. Requires `pip install wandb`.
|
||||
|
||||
## Code Conventions
|
||||
|
||||
- **Line length:** 100 chars (ruff enforced)
|
||||
|
|
@ -81,7 +89,7 @@ soup train --config soup.yaml
|
|||
|
||||
## Tests
|
||||
|
||||
Test suite (~147 tests) lives in `tests/`:
|
||||
Test suite (~184 tests) lives in `tests/`:
|
||||
|
||||
| File | Covers |
|
||||
|---|---|
|
||||
|
|
@ -102,3 +110,6 @@ Test suite (~147 tests) lives in `tests/`:
|
|||
| `test_loader.py` | Data loading (JSONL/JSON/CSV, edge cases) |
|
||||
| `test_validator.py` | `validate_and_stats`, `extended_stats`, `_percentile` |
|
||||
| `test_formats.py` | Reverse conversion, round-trips, edge cases |
|
||||
| `test_merge.py` | Merge command, adapter detection, validation |
|
||||
| `test_export.py` | Export command, GGUF quant types, validation |
|
||||
| `test_resume.py` | Resume checkpoint resolution, W&B flag |
|
||||
|
|
|
|||
85
README.md
85
README.md
|
|
@ -21,7 +21,7 @@
|
|||
<a href="https://pypi.org/project/soup-cli/"><img src="https://img.shields.io/pypi/v/soup-cli?color=blue" alt="PyPI"></a>
|
||||
<img src="https://img.shields.io/badge/python-3.9%2B-blue" alt="Python 3.9+">
|
||||
<img src="https://img.shields.io/badge/license-MIT-green" alt="MIT License">
|
||||
<img src="https://img.shields.io/badge/tests-147%20passed-brightgreen" alt="Tests">
|
||||
<img src="https://img.shields.io/badge/tests-184%20passed-brightgreen" alt="Tests">
|
||||
<a href="https://github.com/MakazhanAlpamys/Soup/actions"><img src="https://github.com/MakazhanAlpamys/Soup/actions/workflows/ci.yml/badge.svg" alt="CI"></a>
|
||||
</p>
|
||||
|
||||
|
|
@ -88,6 +88,16 @@ soup chat --model ./output
|
|||
soup push --model ./output --repo your-username/my-model
|
||||
```
|
||||
|
||||
### 6. Merge & Export
|
||||
|
||||
```bash
|
||||
# Merge LoRA adapter with base model
|
||||
soup merge --adapter ./output
|
||||
|
||||
# Export to GGUF for Ollama / llama.cpp
|
||||
soup export --model ./output --format gguf --quant q4_k_m
|
||||
```
|
||||
|
||||
## Config Example
|
||||
|
||||
```yaml
|
||||
|
|
@ -155,6 +165,69 @@ soup push --model ./output --repo your-username/my-model
|
|||
soup push --model ./output --repo your-username/my-model --private
|
||||
```
|
||||
|
||||
## Merge LoRA Adapter
|
||||
|
||||
Merge a LoRA adapter with its base model into a standalone model:
|
||||
|
||||
```bash
|
||||
# Auto-detect base model from adapter_config.json
|
||||
soup merge --adapter ./output --output ./merged
|
||||
|
||||
# Specify base model and dtype
|
||||
soup merge --adapter ./output --base meta-llama/Llama-3.1-8B --dtype bfloat16
|
||||
```
|
||||
|
||||
## Export to GGUF
|
||||
|
||||
Export models to GGUF format for use with [Ollama](https://ollama.com/) and [llama.cpp](https://github.com/ggerganov/llama.cpp):
|
||||
|
||||
```bash
|
||||
# Export LoRA adapter (auto-merges with base, then converts)
|
||||
soup export --model ./output --format gguf --quant q4_k_m
|
||||
|
||||
# Export with different quantizations
|
||||
soup export --model ./output --format gguf --quant q8_0
|
||||
soup export --model ./output --format gguf --quant f16
|
||||
|
||||
# Export a full (already merged) model
|
||||
soup export --model ./merged --format gguf
|
||||
|
||||
# Specify llama.cpp path manually
|
||||
soup export --model ./output --format gguf --llama-cpp /path/to/llama.cpp
|
||||
```
|
||||
|
||||
Supported quantizations: `q4_0`, `q4_k_m`, `q5_k_m`, `q8_0`, `f16`, `f32`
|
||||
|
||||
After export, use with Ollama:
|
||||
```bash
|
||||
echo 'FROM ./my-model.q4_k_m.gguf' > Modelfile
|
||||
ollama create my-model -f Modelfile
|
||||
ollama run my-model
|
||||
```
|
||||
|
||||
## Resume Training
|
||||
|
||||
Resume a training run from a checkpoint:
|
||||
|
||||
```bash
|
||||
# Auto-detect latest checkpoint in output directory
|
||||
soup train --config soup.yaml --resume auto
|
||||
|
||||
# Resume from a specific checkpoint
|
||||
soup train --config soup.yaml --resume ./output/checkpoint-500
|
||||
```
|
||||
|
||||
## Weights & Biases Integration
|
||||
|
||||
Send training metrics to [W&B](https://wandb.ai/) for cloud-based experiment tracking:
|
||||
|
||||
```bash
|
||||
# Enable W&B logging (requires: pip install wandb)
|
||||
soup train --config soup.yaml --wandb
|
||||
```
|
||||
|
||||
Make sure `WANDB_API_KEY` is set or run `wandb login` first.
|
||||
|
||||
## Data Formats
|
||||
|
||||
Soup supports these formats (auto-detected):
|
||||
|
|
@ -248,6 +321,10 @@ soup eval --model ./output --benchmarks mmlu --run-id run_20260223_143052_a1b2
|
|||
| HuggingFace datasets support | ✅ |
|
||||
| Interactive model chat | ✅ |
|
||||
| Push to HuggingFace Hub | ✅ |
|
||||
| LoRA merge (adapter + base → full model) | ✅ |
|
||||
| Export to GGUF (Ollama / llama.cpp) | ✅ |
|
||||
| Resume training from checkpoint | ✅ |
|
||||
| Weights & Biases integration | ✅ |
|
||||
| Experiment tracking (SQLite) | ✅ |
|
||||
| Data tools (convert, merge, dedup, stats) | ✅ |
|
||||
| Model evaluation (lm-eval) | ✅ |
|
||||
|
|
@ -259,8 +336,12 @@ soup eval --model ./output --benchmarks mmlu --run-id run_20260223_143052_a1b2
|
|||
```
|
||||
soup init [--template chat|code|medical] Create soup.yaml config
|
||||
soup train --config soup.yaml [--dry-run] Start training
|
||||
soup train --resume auto Resume from last checkpoint
|
||||
soup train --wandb Train with W&B logging
|
||||
soup chat --model ./output Interactive chat with model
|
||||
soup push --model ./output --repo user/name Upload to HuggingFace Hub
|
||||
soup merge --adapter ./output Merge LoRA with base model
|
||||
soup export --model ./output --format gguf Export to GGUF (Ollama)
|
||||
soup data inspect <path> View dataset stats
|
||||
soup data validate <path> --format alpaca Check format
|
||||
soup data convert <path> --to chatml Convert between formats
|
||||
|
|
@ -291,7 +372,7 @@ pip install -e ".[dev]"
|
|||
# Lint
|
||||
ruff check soup_cli/ tests/
|
||||
|
||||
# Run unit tests (fast, no GPU needed — 147 tests)
|
||||
# Run unit tests (fast, no GPU needed — 184 tests)
|
||||
pytest tests/ -v
|
||||
|
||||
# Run smoke tests (downloads tiny model, runs real training)
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ build-backend = "hatchling.build"
|
|||
|
||||
[project]
|
||||
name = "soup-cli"
|
||||
version = "0.1.0"
|
||||
version = "0.2.0"
|
||||
description = "Fine-tune LLMs in one command. No SSH, no config hell."
|
||||
readme = "README.md"
|
||||
license = "MIT"
|
||||
|
|
@ -40,6 +40,7 @@ dependencies = [
|
|||
[project.optional-dependencies]
|
||||
eval = ["lm-eval>=0.4.0"]
|
||||
data = ["datasketch>=1.6.0"]
|
||||
wandb = ["wandb>=0.15.0"]
|
||||
dev = ["pytest>=7.0", "ruff>=0.1.0", "pytest-cov>=4.0"]
|
||||
ui = ["fastapi>=0.104.0", "uvicorn>=0.24.0"]
|
||||
|
||||
|
|
|
|||
|
|
@ -1,3 +1,3 @@
|
|||
"""Soup CLI — Fine-tune LLMs in one command."""
|
||||
|
||||
__version__ = "0.1.0"
|
||||
__version__ = "0.2.0"
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ import typer
|
|||
from rich.console import Console
|
||||
|
||||
from soup_cli import __version__
|
||||
from soup_cli.commands import chat, data, eval, init, push, runs, train
|
||||
from soup_cli.commands import chat, data, eval, export, init, merge, push, runs, train
|
||||
|
||||
console = Console()
|
||||
|
||||
|
|
@ -20,6 +20,8 @@ app.command()(init.init)
|
|||
app.command()(train.train)
|
||||
app.command()(chat.chat)
|
||||
app.command()(push.push)
|
||||
app.command(name="export")(export.export)
|
||||
app.command()(merge.merge)
|
||||
app.add_typer(
|
||||
data.app, name="data",
|
||||
help="Dataset tools: inspect, convert, merge, dedup, validate, stats.",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,346 @@
|
|||
"""soup export — convert a model to GGUF format for Ollama / llama.cpp."""
|
||||
|
||||
import json
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import typer
|
||||
from rich.console import Console
|
||||
from rich.panel import Panel
|
||||
|
||||
console = Console()
|
||||
|
||||
SUPPORTED_FORMATS = ("gguf",)
|
||||
GGUF_QUANT_TYPES = ("q4_0", "q4_k_m", "q5_k_m", "q8_0", "f16", "f32")
|
||||
LLAMA_CPP_DIR_NAME = "llama.cpp"
|
||||
|
||||
|
||||
def export(
|
||||
model: str = typer.Option(
|
||||
...,
|
||||
"--model",
|
||||
"-m",
|
||||
help="Path to model directory (full model or LoRA adapter)",
|
||||
),
|
||||
fmt: str = typer.Option(
|
||||
"gguf",
|
||||
"--format",
|
||||
"-f",
|
||||
help="Export format: gguf",
|
||||
),
|
||||
quant: str = typer.Option(
|
||||
"q4_k_m",
|
||||
"--quant",
|
||||
"-q",
|
||||
help="Quantization type: q4_0, q4_k_m, q5_k_m, q8_0, f16, f32",
|
||||
),
|
||||
output: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--output",
|
||||
"-o",
|
||||
help="Output file path. Default: <model-name>.<quant>.gguf",
|
||||
),
|
||||
base: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--base",
|
||||
"-b",
|
||||
help="Base model ID (for LoRA adapters). Auto-detected if not set.",
|
||||
),
|
||||
llama_cpp_path: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--llama-cpp",
|
||||
help="Path to llama.cpp directory. Auto-detected or cloned to ~/.soup/llama.cpp",
|
||||
),
|
||||
):
|
||||
"""Export a model to GGUF format for use with Ollama / llama.cpp."""
|
||||
model_path = Path(model)
|
||||
|
||||
# --- Validate ---
|
||||
if not model_path.exists():
|
||||
console.print(f"[red]Model path not found: {model_path}[/]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
if fmt not in SUPPORTED_FORMATS:
|
||||
console.print(
|
||||
f"[red]Unsupported format: {fmt}[/]\n"
|
||||
f"Supported: {', '.join(SUPPORTED_FORMATS)}"
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
if quant not in GGUF_QUANT_TYPES:
|
||||
console.print(
|
||||
f"[red]Unsupported quantization: {quant}[/]\n"
|
||||
f"Supported: {', '.join(GGUF_QUANT_TYPES)}"
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
# --- Check if LoRA adapter (needs merge first) ---
|
||||
adapter_config_path = model_path / "adapter_config.json"
|
||||
is_adapter = adapter_config_path.exists()
|
||||
merge_dir = None
|
||||
|
||||
if is_adapter:
|
||||
console.print("[yellow]LoRA adapter detected — merging with base model first...[/]")
|
||||
base_model = base or _detect_base_model(adapter_config_path)
|
||||
if not base_model:
|
||||
console.print(
|
||||
"[red]Cannot detect base model from adapter_config.json.[/]\n"
|
||||
"Please specify with [bold]--base[/] flag."
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
merge_dir = model_path.parent / f".soup_merge_tmp_{model_path.name}"
|
||||
_merge_adapter(str(model_path), base_model, str(merge_dir))
|
||||
model_path = merge_dir
|
||||
|
||||
# --- Find llama.cpp ---
|
||||
llama_dir = _find_llama_cpp(llama_cpp_path)
|
||||
|
||||
# --- Convert to GGUF ---
|
||||
model_name = Path(model).name
|
||||
if output:
|
||||
output_path = Path(output)
|
||||
else:
|
||||
output_path = Path(model).parent / f"{model_name}.{quant}.gguf"
|
||||
|
||||
console.print(
|
||||
Panel(
|
||||
f"Model: [bold]{model_path}[/]\n"
|
||||
f"Format: [bold]{fmt}[/]\n"
|
||||
f"Quant: [bold]{quant}[/]\n"
|
||||
f"Output: [bold]{output_path}[/]",
|
||||
title="Export Plan",
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
# Step 1: Convert HF model to GGUF (f16)
|
||||
convert_script = llama_dir / "convert_hf_to_gguf.py"
|
||||
if not convert_script.exists():
|
||||
console.print(
|
||||
f"[red]convert_hf_to_gguf.py not found in {llama_dir}[/]\n"
|
||||
"Make sure llama.cpp is properly cloned."
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
if quant in ("f16", "f32"):
|
||||
# Direct conversion without quantization
|
||||
outtype = "f16" if quant == "f16" else "f32"
|
||||
console.print(f"[dim]Converting to GGUF ({outtype})...[/]")
|
||||
_run_convert(convert_script, model_path, output_path, outtype)
|
||||
else:
|
||||
# Convert to f16 first, then quantize
|
||||
f16_path = output_path.parent / f"{model_name}.f16.gguf"
|
||||
console.print("[dim]Converting to GGUF (f16)...[/]")
|
||||
_run_convert(convert_script, model_path, f16_path, "f16")
|
||||
|
||||
# Quantize
|
||||
console.print(f"[dim]Quantizing to {quant}...[/]")
|
||||
_run_quantize(llama_dir, f16_path, output_path, quant)
|
||||
|
||||
# Clean up intermediate f16 file
|
||||
if f16_path.exists() and f16_path != output_path:
|
||||
f16_path.unlink()
|
||||
|
||||
finally:
|
||||
# Clean up temporary merge directory
|
||||
if merge_dir and merge_dir.exists():
|
||||
console.print("[dim]Cleaning up temporary merge files...[/]")
|
||||
shutil.rmtree(merge_dir, ignore_errors=True)
|
||||
|
||||
if not output_path.exists():
|
||||
console.print("[red]Export failed — output file not created.[/]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
file_size = output_path.stat().st_size
|
||||
size_str = _format_size(file_size)
|
||||
|
||||
console.print(
|
||||
Panel(
|
||||
f"Output: [bold]{output_path}[/]\n"
|
||||
f"Size: [bold]{size_str}[/]\n"
|
||||
f"Quant: [bold]{quant}[/]\n\n"
|
||||
f"Use with Ollama:\n"
|
||||
f" 1. Create a Modelfile:\n"
|
||||
f" [bold]echo 'FROM {output_path}' > Modelfile[/]\n"
|
||||
f" 2. Create the model:\n"
|
||||
f" [bold]ollama create {model_name} -f Modelfile[/]\n"
|
||||
f" 3. Run it:\n"
|
||||
f" [bold]ollama run {model_name}[/]",
|
||||
title="[bold green]Export Complete![/]",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _detect_base_model(adapter_config_path: Path) -> Optional[str]:
|
||||
"""Read base_model_name_or_path from adapter_config.json."""
|
||||
try:
|
||||
with open(adapter_config_path, encoding="utf-8") as f:
|
||||
config = json.load(f)
|
||||
return config.get("base_model_name_or_path")
|
||||
except (json.JSONDecodeError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def _merge_adapter(adapter_path: str, base_model: str, output_dir: str):
|
||||
"""Merge LoRA adapter with base model."""
|
||||
import torch
|
||||
from peft import PeftModel
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
console.print(f"[dim]Loading base model: {base_model}...[/]")
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
base_model,
|
||||
torch_dtype=torch.float16,
|
||||
trust_remote_code=True,
|
||||
device_map="cpu",
|
||||
)
|
||||
|
||||
console.print(f"[dim]Loading LoRA adapter: {adapter_path}...[/]")
|
||||
model = PeftModel.from_pretrained(model, adapter_path)
|
||||
|
||||
console.print("[dim]Merging weights...[/]")
|
||||
model = model.merge_and_unload()
|
||||
|
||||
out = Path(output_dir)
|
||||
out.mkdir(parents=True, exist_ok=True)
|
||||
model.save_pretrained(str(out))
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(adapter_path, trust_remote_code=True)
|
||||
tokenizer.save_pretrained(str(out))
|
||||
console.print("[green]Adapter merged successfully.[/]")
|
||||
|
||||
|
||||
def _find_llama_cpp(user_path: Optional[str] = None) -> Path:
|
||||
"""Find or clone llama.cpp directory."""
|
||||
from soup_cli.utils.constants import SOUP_DIR
|
||||
|
||||
# 1. User-specified path
|
||||
if user_path:
|
||||
path = Path(user_path)
|
||||
if path.exists():
|
||||
return path
|
||||
console.print(f"[red]llama.cpp not found at: {path}[/]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
# 2. LLAMA_CPP_PATH env var
|
||||
import os
|
||||
|
||||
env_path = os.environ.get("LLAMA_CPP_PATH")
|
||||
if env_path:
|
||||
path = Path(env_path)
|
||||
if path.exists():
|
||||
return path
|
||||
|
||||
# 3. Check ~/.soup/llama.cpp
|
||||
soup_llama = Path(SOUP_DIR) / LLAMA_CPP_DIR_NAME
|
||||
if soup_llama.exists() and (soup_llama / "convert_hf_to_gguf.py").exists():
|
||||
return soup_llama
|
||||
|
||||
# 4. Auto-clone
|
||||
console.print("[yellow]llama.cpp not found. Cloning to ~/.soup/llama.cpp...[/]")
|
||||
console.print("[dim]This is a one-time setup for GGUF export.[/]")
|
||||
|
||||
try:
|
||||
subprocess.run(
|
||||
["git", "clone", "--depth", "1", "https://github.com/ggerganov/llama.cpp.git",
|
||||
str(soup_llama)],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
# Install Python requirements for the convert script
|
||||
requirements = soup_llama / "requirements.txt"
|
||||
if requirements.exists():
|
||||
subprocess.run(
|
||||
[sys.executable, "-m", "pip", "install", "-r", str(requirements), "-q"],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
console.print("[green]llama.cpp cloned successfully.[/]")
|
||||
return soup_llama
|
||||
except subprocess.CalledProcessError as exc:
|
||||
console.print(f"[red]Failed to clone llama.cpp: {exc.stderr}[/]")
|
||||
console.print(
|
||||
"Please clone manually:\n"
|
||||
f" [bold]git clone https://github.com/ggerganov/llama.cpp.git {soup_llama}[/]\n"
|
||||
"Or specify path: [bold]--llama-cpp /path/to/llama.cpp[/]"
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
except FileNotFoundError:
|
||||
console.print(
|
||||
"[red]git not found.[/] Please install git or clone llama.cpp manually:\n"
|
||||
f" [bold]git clone https://github.com/ggerganov/llama.cpp.git {soup_llama}[/]"
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
|
||||
def _run_convert(script: Path, model_dir: Path, output_path: Path, outtype: str):
|
||||
"""Run llama.cpp convert_hf_to_gguf.py script."""
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(script),
|
||||
str(model_dir),
|
||||
"--outfile", str(output_path),
|
||||
"--outtype", outtype,
|
||||
]
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode != 0:
|
||||
console.print(f"[red]Conversion failed:[/]\n{result.stderr}")
|
||||
raise typer.Exit(1)
|
||||
|
||||
|
||||
def _run_quantize(llama_dir: Path, input_path: Path, output_path: Path, quant_type: str):
|
||||
"""Run llama-quantize (or llama.cpp/build/bin/llama-quantize)."""
|
||||
# Try to find the quantize binary
|
||||
quantize_bin = _find_quantize_binary(llama_dir)
|
||||
if not quantize_bin:
|
||||
console.print(
|
||||
"[red]llama-quantize binary not found.[/]\n"
|
||||
"Build llama.cpp first:\n"
|
||||
f" [bold]cd {llama_dir} && make llama-quantize[/]\n"
|
||||
"Or use [bold]--quant f16[/] to skip quantization."
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
cmd = [str(quantize_bin), str(input_path), str(output_path), quant_type.upper()]
|
||||
result = subprocess.run(cmd, capture_output=True, text=True)
|
||||
if result.returncode != 0:
|
||||
console.print(f"[red]Quantization failed:[/]\n{result.stderr}")
|
||||
raise typer.Exit(1)
|
||||
|
||||
|
||||
def _find_quantize_binary(llama_dir: Path) -> Optional[Path]:
|
||||
"""Find the llama-quantize binary."""
|
||||
# Check common locations
|
||||
candidates = [
|
||||
llama_dir / "build" / "bin" / "llama-quantize",
|
||||
llama_dir / "build" / "bin" / "llama-quantize.exe",
|
||||
llama_dir / "llama-quantize",
|
||||
llama_dir / "llama-quantize.exe",
|
||||
llama_dir / "build" / "llama-quantize",
|
||||
]
|
||||
for candidate in candidates:
|
||||
if candidate.exists():
|
||||
return candidate
|
||||
|
||||
# Check if it's in PATH
|
||||
which_result = shutil.which("llama-quantize")
|
||||
if which_result:
|
||||
return Path(which_result)
|
||||
|
||||
return None
|
||||
|
||||
|
||||
def _format_size(size_bytes: int) -> str:
|
||||
"""Format bytes into human-readable string."""
|
||||
for unit in ("B", "KB", "MB", "GB"):
|
||||
if size_bytes < 1024:
|
||||
return f"{size_bytes:.1f} {unit}"
|
||||
size_bytes /= 1024
|
||||
return f"{size_bytes:.1f} TB"
|
||||
|
|
@ -0,0 +1,159 @@
|
|||
"""soup merge — merge LoRA adapter with base model into a full model."""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import typer
|
||||
from rich.console import Console
|
||||
from rich.panel import Panel
|
||||
|
||||
console = Console()
|
||||
|
||||
|
||||
def merge(
|
||||
adapter: str = typer.Option(
|
||||
...,
|
||||
"--adapter",
|
||||
"-a",
|
||||
help="Path to the LoRA adapter directory",
|
||||
),
|
||||
base: Optional[str] = typer.Option(
|
||||
None,
|
||||
"--base",
|
||||
"-b",
|
||||
help="Base model ID. Auto-detected from adapter_config.json if not set.",
|
||||
),
|
||||
output: str = typer.Option(
|
||||
"./merged",
|
||||
"--output",
|
||||
"-o",
|
||||
help="Output directory for the merged model",
|
||||
),
|
||||
dtype: str = typer.Option(
|
||||
"float16",
|
||||
"--dtype",
|
||||
help="Data type for the merged model: float16, bfloat16, float32",
|
||||
),
|
||||
):
|
||||
"""Merge a LoRA adapter with its base model into a full model."""
|
||||
adapter_path = Path(adapter)
|
||||
|
||||
# --- Validate adapter ---
|
||||
if not adapter_path.exists():
|
||||
console.print(f"[red]Adapter path not found: {adapter_path}[/]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
adapter_config_path = adapter_path / "adapter_config.json"
|
||||
if not adapter_config_path.exists():
|
||||
console.print(
|
||||
f"[red]Not a LoRA adapter: {adapter_path}[/]\n"
|
||||
"Expected adapter_config.json in the directory."
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
# --- Resolve base model ---
|
||||
if not base:
|
||||
base = _detect_base_model(adapter_config_path)
|
||||
if not base:
|
||||
console.print(
|
||||
"[red]Cannot detect base model from adapter_config.json.[/]\n"
|
||||
"Please specify with [bold]--base[/] flag."
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
# --- Validate dtype ---
|
||||
valid_dtypes = ("float16", "bfloat16", "float32")
|
||||
if dtype not in valid_dtypes:
|
||||
console.print(f"[red]Invalid dtype: {dtype}. Must be one of: {', '.join(valid_dtypes)}[/]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
output_path = Path(output)
|
||||
|
||||
console.print(
|
||||
Panel(
|
||||
f"Adapter: [bold]{adapter_path}[/]\n"
|
||||
f"Base: [bold]{base}[/]\n"
|
||||
f"Output: [bold]{output_path}[/]\n"
|
||||
f"Dtype: [bold]{dtype}[/]",
|
||||
title="Merge Plan",
|
||||
)
|
||||
)
|
||||
|
||||
# --- Merge ---
|
||||
try:
|
||||
import torch
|
||||
from peft import PeftModel
|
||||
from transformers import AutoModelForCausalLM, AutoTokenizer
|
||||
|
||||
dtype_map = {
|
||||
"float16": torch.float16,
|
||||
"bfloat16": torch.bfloat16,
|
||||
"float32": torch.float32,
|
||||
}
|
||||
torch_dtype = dtype_map[dtype]
|
||||
|
||||
console.print(f"[dim]Loading base model: {base}...[/]")
|
||||
model = AutoModelForCausalLM.from_pretrained(
|
||||
base,
|
||||
torch_dtype=torch_dtype,
|
||||
trust_remote_code=True,
|
||||
device_map="cpu",
|
||||
)
|
||||
|
||||
console.print(f"[dim]Loading LoRA adapter: {adapter_path}...[/]")
|
||||
model = PeftModel.from_pretrained(model, str(adapter_path))
|
||||
|
||||
console.print("[dim]Merging weights...[/]")
|
||||
model = model.merge_and_unload()
|
||||
|
||||
console.print(f"[dim]Saving merged model to {output_path}...[/]")
|
||||
output_path.mkdir(parents=True, exist_ok=True)
|
||||
model.save_pretrained(str(output_path))
|
||||
|
||||
console.print("[dim]Saving tokenizer...[/]")
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(adapter_path), trust_remote_code=True)
|
||||
tokenizer.save_pretrained(str(output_path))
|
||||
|
||||
except ImportError as exc:
|
||||
console.print(f"[red]Missing dependency: {exc}[/]")
|
||||
console.print("Run: [bold]pip install torch transformers peft[/]")
|
||||
raise typer.Exit(1)
|
||||
except Exception as exc:
|
||||
console.print(f"[red]Merge failed: {exc}[/]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
# Calculate output size
|
||||
total_size = sum(f.stat().st_size for f in output_path.rglob("*") if f.is_file())
|
||||
size_str = _format_size(total_size)
|
||||
|
||||
console.print(
|
||||
Panel(
|
||||
f"Output: [bold]{output_path}[/]\n"
|
||||
f"Size: [bold]{size_str}[/]\n\n"
|
||||
f"Next steps:\n"
|
||||
f" [bold]soup chat --model {output_path}[/]\n"
|
||||
f" [bold]soup push --model {output_path} --repo user/model[/]\n"
|
||||
f" [bold]soup export --model {output_path} --format gguf[/]",
|
||||
title="[bold green]Merge Complete![/]",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _detect_base_model(adapter_config_path: Path) -> Optional[str]:
|
||||
"""Read base_model_name_or_path from adapter_config.json."""
|
||||
try:
|
||||
with open(adapter_config_path, encoding="utf-8") as f:
|
||||
config = json.load(f)
|
||||
return config.get("base_model_name_or_path")
|
||||
except (json.JSONDecodeError, OSError):
|
||||
return None
|
||||
|
||||
|
||||
def _format_size(size_bytes: int) -> str:
|
||||
"""Format bytes into human-readable string."""
|
||||
for unit in ("B", "KB", "MB", "GB"):
|
||||
if size_bytes < 1024:
|
||||
return f"{size_bytes:.1f} {unit}"
|
||||
size_bytes /= 1024
|
||||
return f"{size_bytes:.1f} TB"
|
||||
|
|
@ -33,6 +33,17 @@ def train(
|
|||
"--dry-run",
|
||||
help="Validate config and data without training",
|
||||
),
|
||||
resume: str = typer.Option(
|
||||
None,
|
||||
"--resume",
|
||||
"-r",
|
||||
help="Resume from checkpoint: path to checkpoint dir, or 'auto' for latest",
|
||||
),
|
||||
wandb: bool = typer.Option(
|
||||
False,
|
||||
"--wandb",
|
||||
help="Enable Weights & Biases logging",
|
||||
),
|
||||
):
|
||||
"""Start training from a soup.yaml config."""
|
||||
config_path = Path(config)
|
||||
|
|
@ -45,6 +56,29 @@ def train(
|
|||
console.print(f"[dim]Loading config from {config_path}...[/]")
|
||||
cfg = load_config(config_path)
|
||||
|
||||
# --- Resolve resume checkpoint (fail fast before heavy operations) ---
|
||||
resume_from = None
|
||||
if resume:
|
||||
resume_from = _resolve_checkpoint(resume, cfg.output, cfg.experiment_name)
|
||||
if resume_from:
|
||||
console.print(f"[green]Resuming from:[/] {resume_from}")
|
||||
else:
|
||||
console.print("[red]No checkpoint found to resume from.[/]")
|
||||
raise typer.Exit(1)
|
||||
|
||||
# --- W&B setup (fail fast if wandb not installed) ---
|
||||
if wandb:
|
||||
try:
|
||||
import wandb as _wandb # noqa: F401
|
||||
|
||||
console.print("[green]W&B logging enabled[/]")
|
||||
except ImportError:
|
||||
console.print(
|
||||
"[red]wandb not installed.[/]\n"
|
||||
"Run: [bold]pip install wandb[/]"
|
||||
)
|
||||
raise typer.Exit(1)
|
||||
|
||||
# Detect hardware
|
||||
device, device_name = detect_device()
|
||||
gpu_info = get_gpu_info()
|
||||
|
|
@ -89,13 +123,14 @@ def train(
|
|||
console.print(f"[dim]Run ID: {run_id}[/]")
|
||||
|
||||
# Build trainer based on task type
|
||||
report_to = "wandb" if wandb else "none"
|
||||
console.print("[dim]Setting up model + trainer...[/]")
|
||||
if cfg.task == "dpo":
|
||||
from soup_cli.trainer.dpo import DPOTrainerWrapper
|
||||
|
||||
trainer_wrapper = DPOTrainerWrapper(cfg, device=device)
|
||||
trainer_wrapper = DPOTrainerWrapper(cfg, device=device, report_to=report_to)
|
||||
else:
|
||||
trainer_wrapper = SFTTrainerWrapper(cfg, device=device)
|
||||
trainer_wrapper = SFTTrainerWrapper(cfg, device=device, report_to=report_to)
|
||||
trainer_wrapper.setup(dataset)
|
||||
|
||||
# Train with live display and experiment tracking
|
||||
|
|
@ -104,7 +139,8 @@ def train(
|
|||
|
||||
try:
|
||||
result = trainer_wrapper.train(
|
||||
display=display, tracker=tracker, run_id=run_id
|
||||
display=display, tracker=tracker, run_id=run_id,
|
||||
resume_from_checkpoint=resume_from,
|
||||
)
|
||||
|
||||
# Save completion to tracker
|
||||
|
|
@ -129,7 +165,38 @@ def train(
|
|||
f"Run ID: [bold]{run_id}[/]\n\n"
|
||||
f"Quick test: [bold]soup chat --model {result['output_dir']}[/]\n"
|
||||
f"Push to HF: [bold]soup push --model {result['output_dir']}[/]\n"
|
||||
f"Merge LoRA: [bold]soup merge --adapter {result['output_dir']}[/]\n"
|
||||
f"Export GGUF: [bold]soup export --model {result['output_dir']}[/]\n"
|
||||
f"Run details: [bold]soup runs show {run_id}[/]",
|
||||
title="[bold green]Training Complete![/]",
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _resolve_checkpoint(resume: str, output_dir: str, experiment_name: str = None) -> str:
|
||||
"""Resolve the checkpoint path from --resume argument.
|
||||
|
||||
If resume == "auto", find the latest checkpoint in the output directory.
|
||||
Otherwise, treat it as a direct path to a checkpoint directory.
|
||||
"""
|
||||
if resume.lower() == "auto":
|
||||
base = Path(output_dir)
|
||||
if experiment_name:
|
||||
base = base / experiment_name
|
||||
|
||||
if not base.exists():
|
||||
return None
|
||||
|
||||
checkpoints = sorted(
|
||||
[d for d in base.iterdir() if d.is_dir() and d.name.startswith("checkpoint-")],
|
||||
key=lambda d: int(d.name.split("-")[-1]) if d.name.split("-")[-1].isdigit() else 0,
|
||||
)
|
||||
if checkpoints:
|
||||
return str(checkpoints[-1])
|
||||
return None
|
||||
|
||||
# Direct path
|
||||
checkpoint_path = Path(resume)
|
||||
if checkpoint_path.exists() and checkpoint_path.is_dir():
|
||||
return str(checkpoint_path)
|
||||
return None
|
||||
|
|
|
|||
|
|
@ -21,9 +21,10 @@ class DPOTrainerWrapper:
|
|||
- rejected: the less preferred response
|
||||
"""
|
||||
|
||||
def __init__(self, config: SoupConfig, device: str = "cuda"):
|
||||
def __init__(self, config: SoupConfig, device: str = "cuda", report_to: str = "none"):
|
||||
self.config = config
|
||||
self.device = device
|
||||
self.report_to = report_to
|
||||
self.model = None
|
||||
self.ref_model = None
|
||||
self.tokenizer = None
|
||||
|
|
@ -142,7 +143,7 @@ class DPOTrainerWrapper:
|
|||
save_steps=tcfg.save_steps,
|
||||
save_total_limit=3,
|
||||
bf16=self.device == "cuda",
|
||||
report_to="none",
|
||||
report_to=self.report_to,
|
||||
remove_unused_columns=False,
|
||||
beta=tcfg.dpo_beta,
|
||||
max_length=cfg.data.max_length,
|
||||
|
|
@ -165,6 +166,7 @@ class DPOTrainerWrapper:
|
|||
display: Optional[object] = None,
|
||||
tracker: Optional[object] = None,
|
||||
run_id: str = "",
|
||||
resume_from_checkpoint: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""Run DPO training and return results summary."""
|
||||
start = time.time()
|
||||
|
|
@ -177,7 +179,7 @@ class DPOTrainerWrapper:
|
|||
SoupTrainerCallback(display, tracker=tracker, run_id=run_id)
|
||||
)
|
||||
|
||||
self.trainer.train()
|
||||
self.trainer.train(resume_from_checkpoint=resume_from_checkpoint)
|
||||
duration = time.time() - start
|
||||
|
||||
# Save final model (LoRA adapter)
|
||||
|
|
|
|||
|
|
@ -15,9 +15,10 @@ console = Console()
|
|||
class SFTTrainerWrapper:
|
||||
"""High-level wrapper that sets up model + tokenizer + trainer from SoupConfig."""
|
||||
|
||||
def __init__(self, config: SoupConfig, device: str = "cuda"):
|
||||
def __init__(self, config: SoupConfig, device: str = "cuda", report_to: str = "none"):
|
||||
self.config = config
|
||||
self.device = device
|
||||
self.report_to = report_to
|
||||
self.model = None
|
||||
self.tokenizer = None
|
||||
self.trainer = None
|
||||
|
|
@ -152,7 +153,7 @@ class SFTTrainerWrapper:
|
|||
save_steps=tcfg.save_steps,
|
||||
save_total_limit=3,
|
||||
bf16=self.device == "cuda",
|
||||
report_to="none",
|
||||
report_to=self.report_to,
|
||||
remove_unused_columns=False,
|
||||
)
|
||||
|
||||
|
|
@ -172,6 +173,7 @@ class SFTTrainerWrapper:
|
|||
display: Optional[object] = None,
|
||||
tracker: Optional[object] = None,
|
||||
run_id: str = "",
|
||||
resume_from_checkpoint: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""Run training and return results summary."""
|
||||
start = time.time()
|
||||
|
|
@ -184,7 +186,7 @@ class SFTTrainerWrapper:
|
|||
SoupTrainerCallback(display, tracker=tracker, run_id=run_id)
|
||||
)
|
||||
|
||||
self.trainer.train()
|
||||
self.trainer.train(resume_from_checkpoint=resume_from_checkpoint)
|
||||
duration = time.time() - start
|
||||
|
||||
# Save final model (LoRA adapter)
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ runner = CliRunner()
|
|||
def test_version():
|
||||
result = runner.invoke(app, ["version"])
|
||||
assert result.exit_code == 0
|
||||
assert "0.1.0" in result.output
|
||||
assert "0.2.0" in result.output
|
||||
|
||||
|
||||
def test_help():
|
||||
|
|
@ -68,3 +68,5 @@ def test_help_shows_all_commands():
|
|||
assert "push" in result.output
|
||||
assert "train" in result.output
|
||||
assert "init" in result.output
|
||||
assert "export" in result.output
|
||||
assert "merge" in result.output
|
||||
|
|
|
|||
|
|
@ -0,0 +1,122 @@
|
|||
"""Tests for soup export command."""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from typer.testing import CliRunner
|
||||
|
||||
from soup_cli.cli import app
|
||||
from soup_cli.commands.export import (
|
||||
GGUF_QUANT_TYPES,
|
||||
SUPPORTED_FORMATS,
|
||||
_detect_base_model,
|
||||
_find_quantize_binary,
|
||||
_format_size,
|
||||
)
|
||||
|
||||
runner = CliRunner()
|
||||
|
||||
|
||||
# --- _format_size ---
|
||||
|
||||
def test_format_size_bytes():
|
||||
assert _format_size(100) == "100.0 B"
|
||||
|
||||
|
||||
def test_format_size_gb():
|
||||
assert _format_size(2 * 1024**3) == "2.0 GB"
|
||||
|
||||
|
||||
# --- _detect_base_model ---
|
||||
|
||||
def test_detect_base_model_valid(tmp_path: Path):
|
||||
config = tmp_path / "adapter_config.json"
|
||||
config.write_text(json.dumps({
|
||||
"base_model_name_or_path": "meta-llama/Llama-3.1-8B",
|
||||
}))
|
||||
assert _detect_base_model(config) == "meta-llama/Llama-3.1-8B"
|
||||
|
||||
|
||||
def test_detect_base_model_no_key(tmp_path: Path):
|
||||
config = tmp_path / "adapter_config.json"
|
||||
config.write_text(json.dumps({"r": 64}))
|
||||
assert _detect_base_model(config) is None
|
||||
|
||||
|
||||
def test_detect_base_model_bad_json(tmp_path: Path):
|
||||
config = tmp_path / "adapter_config.json"
|
||||
config.write_text("bad json")
|
||||
assert _detect_base_model(config) is None
|
||||
|
||||
|
||||
# --- _find_quantize_binary ---
|
||||
|
||||
def test_find_quantize_binary_not_found(tmp_path: Path):
|
||||
"""Should return None if no quantize binary exists."""
|
||||
assert _find_quantize_binary(tmp_path) is None
|
||||
|
||||
|
||||
def test_find_quantize_binary_in_build(tmp_path: Path):
|
||||
"""Should find binary in build/bin/."""
|
||||
bin_dir = tmp_path / "build" / "bin"
|
||||
bin_dir.mkdir(parents=True)
|
||||
quantize = bin_dir / "llama-quantize"
|
||||
quantize.write_text("fake binary")
|
||||
assert _find_quantize_binary(tmp_path) == quantize
|
||||
|
||||
|
||||
# --- Constants ---
|
||||
|
||||
def test_supported_formats():
|
||||
assert "gguf" in SUPPORTED_FORMATS
|
||||
|
||||
|
||||
def test_gguf_quant_types():
|
||||
assert "q4_k_m" in GGUF_QUANT_TYPES
|
||||
assert "f16" in GGUF_QUANT_TYPES
|
||||
assert len(GGUF_QUANT_TYPES) >= 4
|
||||
|
||||
|
||||
# --- CLI validation ---
|
||||
|
||||
def test_export_missing_model():
|
||||
result = runner.invoke(app, ["export", "--model", "/nonexistent"])
|
||||
assert result.exit_code == 1
|
||||
assert "not found" in result.output.lower()
|
||||
|
||||
|
||||
def test_export_unsupported_format(tmp_path: Path):
|
||||
model_dir = tmp_path / "model"
|
||||
model_dir.mkdir()
|
||||
result = runner.invoke(
|
||||
app, ["export", "--model", str(model_dir), "--format", "onnx"]
|
||||
)
|
||||
assert result.exit_code == 1
|
||||
assert "unsupported format" in result.output.lower()
|
||||
|
||||
|
||||
def test_export_unsupported_quant(tmp_path: Path):
|
||||
model_dir = tmp_path / "model"
|
||||
model_dir.mkdir()
|
||||
result = runner.invoke(
|
||||
app, ["export", "--model", str(model_dir), "--quant", "q2_k"]
|
||||
)
|
||||
assert result.exit_code == 1
|
||||
assert "unsupported quantization" in result.output.lower()
|
||||
|
||||
|
||||
def test_export_adapter_no_base(tmp_path: Path):
|
||||
"""LoRA adapter without detectable base model and no --base flag."""
|
||||
adapter_dir = tmp_path / "adapter"
|
||||
adapter_dir.mkdir()
|
||||
(adapter_dir / "adapter_config.json").write_text(json.dumps({"r": 64}))
|
||||
result = runner.invoke(app, ["export", "--model", str(adapter_dir)])
|
||||
assert result.exit_code == 1
|
||||
assert "base" in result.output.lower()
|
||||
|
||||
|
||||
def test_export_help():
|
||||
result = runner.invoke(app, ["export", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "gguf" in result.output.lower()
|
||||
assert "quant" in result.output.lower()
|
||||
|
|
@ -0,0 +1,104 @@
|
|||
"""Tests for soup merge command."""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
from typer.testing import CliRunner
|
||||
|
||||
from soup_cli.cli import app
|
||||
from soup_cli.commands.merge import _detect_base_model, _format_size
|
||||
|
||||
runner = CliRunner()
|
||||
|
||||
|
||||
# --- _format_size ---
|
||||
|
||||
def test_format_size_bytes():
|
||||
assert _format_size(512) == "512.0 B"
|
||||
|
||||
|
||||
def test_format_size_kb():
|
||||
assert _format_size(2048) == "2.0 KB"
|
||||
|
||||
|
||||
def test_format_size_mb():
|
||||
assert _format_size(5 * 1024 * 1024) == "5.0 MB"
|
||||
|
||||
|
||||
def test_format_size_gb():
|
||||
assert _format_size(3 * 1024**3) == "3.0 GB"
|
||||
|
||||
|
||||
# --- _detect_base_model ---
|
||||
|
||||
def test_detect_base_model(tmp_path: Path):
|
||||
config = tmp_path / "adapter_config.json"
|
||||
config.write_text(json.dumps({
|
||||
"base_model_name_or_path": "meta-llama/Llama-3.1-8B",
|
||||
}))
|
||||
assert _detect_base_model(config) == "meta-llama/Llama-3.1-8B"
|
||||
|
||||
|
||||
def test_detect_base_model_missing_key(tmp_path: Path):
|
||||
config = tmp_path / "adapter_config.json"
|
||||
config.write_text(json.dumps({"r": 64}))
|
||||
assert _detect_base_model(config) is None
|
||||
|
||||
|
||||
def test_detect_base_model_bad_json(tmp_path: Path):
|
||||
config = tmp_path / "adapter_config.json"
|
||||
config.write_text("not valid json")
|
||||
assert _detect_base_model(config) is None
|
||||
|
||||
|
||||
def test_detect_base_model_missing_file(tmp_path: Path):
|
||||
config = tmp_path / "nonexistent.json"
|
||||
assert _detect_base_model(config) is None
|
||||
|
||||
|
||||
# --- CLI validation ---
|
||||
|
||||
def test_merge_missing_adapter():
|
||||
result = runner.invoke(app, ["merge", "--adapter", "/nonexistent"])
|
||||
assert result.exit_code == 1
|
||||
assert "not found" in result.output.lower()
|
||||
|
||||
|
||||
def test_merge_not_a_lora_adapter(tmp_path: Path):
|
||||
"""Directory without adapter_config.json should fail."""
|
||||
model_dir = tmp_path / "model"
|
||||
model_dir.mkdir()
|
||||
result = runner.invoke(app, ["merge", "--adapter", str(model_dir)])
|
||||
assert result.exit_code == 1
|
||||
assert "not a lora adapter" in result.output.lower()
|
||||
|
||||
|
||||
def test_merge_no_base_model_detected(tmp_path: Path):
|
||||
"""Adapter with empty config (no base_model_name_or_path) and no --base flag."""
|
||||
adapter_dir = tmp_path / "adapter"
|
||||
adapter_dir.mkdir()
|
||||
(adapter_dir / "adapter_config.json").write_text(json.dumps({"r": 64}))
|
||||
result = runner.invoke(app, ["merge", "--adapter", str(adapter_dir)])
|
||||
assert result.exit_code == 1
|
||||
assert "base" in result.output.lower()
|
||||
|
||||
|
||||
def test_merge_invalid_dtype(tmp_path: Path):
|
||||
"""Invalid dtype should fail."""
|
||||
adapter_dir = tmp_path / "adapter"
|
||||
adapter_dir.mkdir()
|
||||
(adapter_dir / "adapter_config.json").write_text(json.dumps({
|
||||
"base_model_name_or_path": "meta-llama/Llama-3.1-8B",
|
||||
}))
|
||||
result = runner.invoke(
|
||||
app, ["merge", "--adapter", str(adapter_dir), "--dtype", "int8"]
|
||||
)
|
||||
assert result.exit_code == 1
|
||||
assert "invalid dtype" in result.output.lower()
|
||||
|
||||
|
||||
def test_merge_help():
|
||||
result = runner.invoke(app, ["merge", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "adapter" in result.output.lower()
|
||||
assert "base" in result.output.lower()
|
||||
|
|
@ -0,0 +1,107 @@
|
|||
"""Tests for resume training and W&B integration."""
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
from typer.testing import CliRunner
|
||||
|
||||
from soup_cli.cli import app
|
||||
from soup_cli.commands.train import _resolve_checkpoint
|
||||
|
||||
runner = CliRunner()
|
||||
|
||||
|
||||
# --- _resolve_checkpoint ---
|
||||
|
||||
def test_resolve_checkpoint_auto_finds_latest(tmp_path: Path):
|
||||
"""auto should find the latest checkpoint by number."""
|
||||
output_dir = tmp_path / "output"
|
||||
output_dir.mkdir()
|
||||
(output_dir / "checkpoint-100").mkdir()
|
||||
(output_dir / "checkpoint-200").mkdir()
|
||||
(output_dir / "checkpoint-50").mkdir()
|
||||
|
||||
result = _resolve_checkpoint("auto", str(output_dir))
|
||||
assert result == str(output_dir / "checkpoint-200")
|
||||
|
||||
|
||||
def test_resolve_checkpoint_auto_with_experiment_name(tmp_path: Path):
|
||||
"""auto should look inside experiment_name subdirectory."""
|
||||
output_dir = tmp_path / "output"
|
||||
exp_dir = output_dir / "my-experiment"
|
||||
exp_dir.mkdir(parents=True)
|
||||
(exp_dir / "checkpoint-100").mkdir()
|
||||
(exp_dir / "checkpoint-300").mkdir()
|
||||
|
||||
result = _resolve_checkpoint("auto", str(output_dir), experiment_name="my-experiment")
|
||||
assert result == str(exp_dir / "checkpoint-300")
|
||||
|
||||
|
||||
def test_resolve_checkpoint_auto_no_checkpoints(tmp_path: Path):
|
||||
"""auto should return None if no checkpoints exist."""
|
||||
output_dir = tmp_path / "output"
|
||||
output_dir.mkdir()
|
||||
|
||||
result = _resolve_checkpoint("auto", str(output_dir))
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_resolve_checkpoint_auto_missing_dir(tmp_path: Path):
|
||||
"""auto should return None if output dir doesn't exist."""
|
||||
result = _resolve_checkpoint("auto", str(tmp_path / "nonexistent"))
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_resolve_checkpoint_direct_path(tmp_path: Path):
|
||||
"""Direct path to checkpoint should be returned as-is."""
|
||||
checkpoint = tmp_path / "checkpoint-100"
|
||||
checkpoint.mkdir()
|
||||
|
||||
result = _resolve_checkpoint(str(checkpoint), str(tmp_path))
|
||||
assert result == str(checkpoint)
|
||||
|
||||
|
||||
def test_resolve_checkpoint_direct_path_nonexistent():
|
||||
"""Nonexistent direct path should return None."""
|
||||
result = _resolve_checkpoint("/nonexistent/checkpoint-100", "/output")
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_resolve_checkpoint_auto_ignores_non_checkpoint_dirs(tmp_path: Path):
|
||||
"""auto should ignore directories that don't start with checkpoint-."""
|
||||
output_dir = tmp_path / "output"
|
||||
output_dir.mkdir()
|
||||
(output_dir / "checkpoint-100").mkdir()
|
||||
(output_dir / "logs").mkdir()
|
||||
(output_dir / "runs").mkdir()
|
||||
|
||||
result = _resolve_checkpoint("auto", str(output_dir))
|
||||
assert result == str(output_dir / "checkpoint-100")
|
||||
|
||||
|
||||
# --- CLI flags ---
|
||||
|
||||
def test_train_resume_flag_in_help():
|
||||
result = runner.invoke(app, ["train", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "--resume" in result.output
|
||||
|
||||
|
||||
def test_train_wandb_flag_in_help():
|
||||
result = runner.invoke(app, ["train", "--help"])
|
||||
assert result.exit_code == 0
|
||||
assert "--wandb" in result.output
|
||||
|
||||
|
||||
def test_train_resume_nonexistent_checkpoint(tmp_path: Path):
|
||||
"""Resume with nonexistent checkpoint path should fail (after config validation)."""
|
||||
config_file = tmp_path / "soup.yaml"
|
||||
config_file.write_text("""
|
||||
base: meta-llama/Llama-3.1-8B
|
||||
data:
|
||||
train: ./data.jsonl
|
||||
""")
|
||||
result = runner.invoke(
|
||||
app, ["train", "--config", str(config_file), "--resume", "/nonexistent/checkpoint"]
|
||||
)
|
||||
assert result.exit_code == 1
|
||||
assert "no checkpoint found" in result.output.lower()
|
||||
Loading…
Reference in New Issue