refactor
This commit is contained in:
parent
4ef25032c1
commit
22f9a65772
3 changed files with 166 additions and 208 deletions
228
cli.py
228
cli.py
|
|
@ -1,4 +1,3 @@
|
|||
import json
|
||||
import logging
|
||||
import sys
|
||||
import time
|
||||
|
|
@ -6,7 +5,8 @@ from pathlib import Path
|
|||
from typing import Optional, List
|
||||
|
||||
import typer
|
||||
import yaml
|
||||
|
||||
from cli.config import Config, load_config
|
||||
|
||||
app = typer.Typer(
|
||||
help="Command-line interface for Unsloth training, chat, and export.",
|
||||
|
|
@ -23,65 +23,6 @@ def configure_logging(verbose: bool):
|
|||
)
|
||||
|
||||
|
||||
def _load_config(config_path: Optional[Path]) -> dict:
|
||||
if not config_path:
|
||||
return {}
|
||||
path = Path(config_path)
|
||||
if not path.exists():
|
||||
raise typer.BadParameter(f"Config file not found: {config_path}")
|
||||
text = path.read_text(encoding="utf-8")
|
||||
if path.suffix.lower() in {".yaml", ".yml"}:
|
||||
return yaml.safe_load(text) or {}
|
||||
else:
|
||||
return json.loads(text or "{}")
|
||||
|
||||
|
||||
def _flatten_config(cfg: dict) -> dict:
|
||||
"""
|
||||
Flatten nested config sections into a single dict.
|
||||
|
||||
Expected sections:
|
||||
data: dataset, local_dataset, format_type
|
||||
training: training_type, max_seq_length, load_in_4bit, output_dir, etc.
|
||||
lora: lora_r, lora_alpha, lora_dropout, target_modules, etc.
|
||||
vision: finetune_vision_layers, finetune_language_layers, etc.
|
||||
logging: enable_wandb, wandb_project, wandb_token, enable_tensorboard, etc.
|
||||
"""
|
||||
if not isinstance(cfg, dict):
|
||||
return {}
|
||||
|
||||
flattened = {}
|
||||
|
||||
# Handle top-level 'model' key
|
||||
if "model" in cfg:
|
||||
flattened["model"] = cfg["model"]
|
||||
|
||||
sections = ["data", "training", "lora", "vision", "logging"]
|
||||
|
||||
for section in sections:
|
||||
if section in cfg and isinstance(cfg[section], dict):
|
||||
flattened.update(cfg[section])
|
||||
|
||||
return flattened
|
||||
|
||||
|
||||
def _merge_config(cfg: dict, defaults: dict, overrides: dict) -> dict:
|
||||
"""
|
||||
Merge CLI overrides with config and defaults.
|
||||
CLI override wins, then config value, then default.
|
||||
"""
|
||||
merged = {}
|
||||
for key, default in defaults.items():
|
||||
cli_val = overrides.get(key, None)
|
||||
if cli_val is not None:
|
||||
merged[key] = cli_val
|
||||
elif key in cfg and cfg[key] is not None:
|
||||
merged[key] = cfg[key]
|
||||
else:
|
||||
merged[key] = default
|
||||
return merged
|
||||
|
||||
|
||||
@app.command()
|
||||
def train(
|
||||
model: Optional[str] = typer.Option(
|
||||
|
|
@ -198,99 +139,22 @@ def train(
|
|||
"""
|
||||
Launch training using the existing Unsloth training backend.
|
||||
"""
|
||||
cfg = _load_config(config)
|
||||
cfg = _flatten_config(cfg)
|
||||
try:
|
||||
cfg = load_config(config)
|
||||
except FileNotFoundError as e:
|
||||
typer.echo(f"Error: {e}", err=True)
|
||||
raise typer.Exit(code=2)
|
||||
|
||||
# Defaults (match previous behavior)
|
||||
defaults = {
|
||||
"model": None,
|
||||
"training_type": "lora",
|
||||
"max_seq_length": 2048,
|
||||
"load_in_4bit": True,
|
||||
"output_dir": Path("./outputs"),
|
||||
"dataset": None,
|
||||
"local_dataset": None,
|
||||
"format_type": "auto",
|
||||
"num_epochs": 3,
|
||||
"learning_rate": 2e-4,
|
||||
"batch_size": 2,
|
||||
"gradient_accumulation_steps": 4,
|
||||
"warmup_steps": 5,
|
||||
"max_steps": 0,
|
||||
"save_steps": 0,
|
||||
"weight_decay": 0.01,
|
||||
"random_seed": 3407,
|
||||
"packing": False,
|
||||
"train_on_completions": False,
|
||||
"lora_r": 64,
|
||||
"lora_alpha": 16,
|
||||
"lora_dropout": 0.0,
|
||||
"gradient_checkpointing": True,
|
||||
"target_modules": "q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj",
|
||||
"vision_all_linear": False,
|
||||
"finetune_vision_layers": True,
|
||||
"finetune_language_layers": True,
|
||||
"finetune_attention_modules": True,
|
||||
"finetune_mlp_modules": True,
|
||||
"use_rslora": False,
|
||||
"use_loftq": False,
|
||||
"enable_wandb": False,
|
||||
"wandb_project": "unsloth-training",
|
||||
"enable_tensorboard": False,
|
||||
"tensorboard_dir": "runs",
|
||||
}
|
||||
# Apply CLI overrides
|
||||
cli_args = {k: v for k, v in locals().items() if k not in ("config", "verbose", "hf_token", "cfg")}
|
||||
cfg.apply_overrides(**cli_args)
|
||||
|
||||
overrides = {
|
||||
"training_type": training_type,
|
||||
"max_seq_length": max_seq_length,
|
||||
"load_in_4bit": load_in_4bit,
|
||||
"output_dir": output_dir,
|
||||
"dataset": dataset,
|
||||
"local_dataset": local_dataset,
|
||||
"format_type": format_type,
|
||||
"num_epochs": num_epochs,
|
||||
"learning_rate": learning_rate,
|
||||
"batch_size": batch_size,
|
||||
"gradient_accumulation_steps": gradient_accumulation_steps,
|
||||
"warmup_steps": warmup_steps,
|
||||
"max_steps": max_steps,
|
||||
"save_steps": save_steps,
|
||||
"weight_decay": weight_decay,
|
||||
"random_seed": random_seed,
|
||||
"packing": packing,
|
||||
"train_on_completions": train_on_completions,
|
||||
"lora_r": lora_r,
|
||||
"lora_alpha": lora_alpha,
|
||||
"lora_dropout": lora_dropout,
|
||||
"gradient_checkpointing": gradient_checkpointing,
|
||||
"target_modules": target_modules,
|
||||
"vision_all_linear": vision_all_linear,
|
||||
"finetune_vision_layers": finetune_vision_layers,
|
||||
"finetune_language_layers": finetune_language_layers,
|
||||
"finetune_attention_modules": finetune_attention_modules,
|
||||
"finetune_mlp_modules": finetune_mlp_modules,
|
||||
"use_rslora": use_rslora,
|
||||
"use_loftq": use_loftq,
|
||||
"enable_wandb": enable_wandb,
|
||||
"wandb_project": wandb_project,
|
||||
"wandb_token": wandb_token,
|
||||
"enable_tensorboard": enable_tensorboard,
|
||||
"tensorboard_dir": tensorboard_dir,
|
||||
}
|
||||
|
||||
merged = _merge_config(cfg, defaults, overrides)
|
||||
|
||||
model_val = merged.get("model")
|
||||
if not model_val:
|
||||
# Validate required fields
|
||||
if not cfg.model:
|
||||
typer.echo("Error: provide --model or set model in --config", err=True)
|
||||
raise typer.Exit(code=2)
|
||||
|
||||
# Convert specific types
|
||||
output_dir_val = Path(merged["output_dir"])
|
||||
dataset_val = merged.get("dataset")
|
||||
local_dataset_val = merged.get("local_dataset")
|
||||
|
||||
if not dataset_val and not local_dataset_val:
|
||||
if not cfg.data.dataset and not cfg.data.local_dataset:
|
||||
typer.echo(
|
||||
"Error: provide --dataset or --local-dataset (or via --config)", err=True
|
||||
)
|
||||
|
|
@ -304,86 +168,38 @@ def train(
|
|||
trainer = UnslothTrainer()
|
||||
|
||||
model_config = ModelConfig.from_ui_selection(
|
||||
dropdown_value=model_val, search_value=None, hf_token=hf_token, is_lora=False
|
||||
dropdown_value=cfg.model, search_value=None, hf_token=hf_token, is_lora=False
|
||||
)
|
||||
if not model_config:
|
||||
typer.echo("Could not resolve model config", err=True)
|
||||
raise typer.Exit(code=1)
|
||||
|
||||
is_vision = model_config.is_vision
|
||||
use_lora = cfg.training.training_type.lower() == "lora"
|
||||
|
||||
if not trainer.load_model(
|
||||
model_name=model_config.identifier,
|
||||
max_seq_length=merged["max_seq_length"],
|
||||
load_in_4bit=merged["load_in_4bit"]
|
||||
if merged["training_type"].lower() == "lora"
|
||||
else False,
|
||||
max_seq_length=cfg.training.max_seq_length,
|
||||
load_in_4bit=cfg.training.load_in_4bit if use_lora else False,
|
||||
hf_token=hf_token,
|
||||
):
|
||||
typer.echo("Model load failed", err=True)
|
||||
raise typer.Exit(code=1)
|
||||
|
||||
use_lora = merged["training_type"].lower() == "lora"
|
||||
|
||||
# Match UI behavior for target modules:
|
||||
# - Text: use parsed target modules list
|
||||
# - Vision: if vision_all_linear, use ["all-linear"]; otherwise empty list
|
||||
target_modules_list = [
|
||||
m.strip() for m in merged["target_modules"].split(",") if m.strip()
|
||||
]
|
||||
if use_lora and is_vision:
|
||||
if merged["vision_all_linear"]:
|
||||
target_modules_list = ["all-linear"]
|
||||
else:
|
||||
target_modules_list = []
|
||||
|
||||
if not trainer.prepare_model_for_training(
|
||||
use_lora=use_lora,
|
||||
finetune_vision_layers=merged["finetune_vision_layers"],
|
||||
finetune_language_layers=merged["finetune_language_layers"],
|
||||
finetune_attention_modules=merged["finetune_attention_modules"],
|
||||
finetune_mlp_modules=merged["finetune_mlp_modules"],
|
||||
target_modules=target_modules_list,
|
||||
lora_r=merged["lora_r"],
|
||||
lora_alpha=merged["lora_alpha"],
|
||||
lora_dropout=merged["lora_dropout"],
|
||||
use_gradient_checkpointing=merged["gradient_checkpointing"],
|
||||
use_rslora=merged["use_rslora"],
|
||||
use_loftq=merged["use_loftq"],
|
||||
):
|
||||
if not trainer.prepare_model_for_training(**cfg.model_kwargs(use_lora, is_vision)):
|
||||
typer.echo("Model preparation failed", err=True)
|
||||
raise typer.Exit(code=1)
|
||||
|
||||
ds = trainer.load_and_format_dataset(
|
||||
dataset_source=dataset_val or "",
|
||||
format_type=merged["format_type"],
|
||||
local_datasets=local_dataset_val,
|
||||
dataset_source=cfg.data.dataset or "",
|
||||
format_type=cfg.data.format_type,
|
||||
local_datasets=cfg.data.local_dataset,
|
||||
)
|
||||
if ds is None:
|
||||
typer.echo("Dataset load failed", err=True)
|
||||
raise typer.Exit(code=1)
|
||||
|
||||
started = trainer.start_training(
|
||||
dataset=ds,
|
||||
output_dir=str(output_dir_val),
|
||||
num_epochs=merged["num_epochs"],
|
||||
learning_rate=merged["learning_rate"],
|
||||
batch_size=merged["batch_size"],
|
||||
gradient_accumulation_steps=merged["gradient_accumulation_steps"],
|
||||
warmup_steps=merged["warmup_steps"],
|
||||
max_steps=merged["max_steps"],
|
||||
save_steps=merged["save_steps"],
|
||||
weight_decay=merged["weight_decay"],
|
||||
random_seed=merged["random_seed"],
|
||||
packing=merged["packing"],
|
||||
train_on_completions=merged["train_on_completions"],
|
||||
enable_wandb=merged["enable_wandb"],
|
||||
wandb_project=merged["wandb_project"],
|
||||
wandb_token=merged.get("wandb_token"),
|
||||
enable_tensorboard=merged["enable_tensorboard"],
|
||||
tensorboard_dir=merged["tensorboard_dir"],
|
||||
max_seq_length=merged["max_seq_length"],
|
||||
)
|
||||
started = trainer.start_training(dataset=ds, **cfg.training_kwargs())
|
||||
|
||||
if not started:
|
||||
typer.echo("Training failed to start", err=True)
|
||||
|
|
|
|||
0
cli/__init__.py
Normal file
0
cli/__init__.py
Normal file
142
cli/config.py
Normal file
142
cli/config.py
Normal file
|
|
@ -0,0 +1,142 @@
|
|||
from pathlib import Path
|
||||
from typing import Optional, List
|
||||
|
||||
import yaml
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class DataConfig(BaseModel):
|
||||
dataset: Optional[str] = None
|
||||
local_dataset: Optional[List[str]] = None
|
||||
format_type: str = "auto"
|
||||
|
||||
|
||||
class TrainingConfig(BaseModel):
|
||||
training_type: str = "lora"
|
||||
max_seq_length: int = 2048
|
||||
load_in_4bit: bool = True
|
||||
output_dir: Path = Path("./outputs")
|
||||
num_epochs: int = 3
|
||||
learning_rate: float = 2e-4
|
||||
batch_size: int = 2
|
||||
gradient_accumulation_steps: int = 4
|
||||
warmup_steps: int = 5
|
||||
max_steps: int = 0
|
||||
save_steps: int = 0
|
||||
weight_decay: float = 0.01
|
||||
random_seed: int = 3407
|
||||
packing: bool = False
|
||||
train_on_completions: bool = False
|
||||
gradient_checkpointing: bool = True
|
||||
|
||||
|
||||
class LoraConfig(BaseModel):
|
||||
lora_r: int = 64
|
||||
lora_alpha: int = 16
|
||||
lora_dropout: float = 0.0
|
||||
target_modules: str = "q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj"
|
||||
vision_all_linear: bool = False
|
||||
use_rslora: bool = False
|
||||
use_loftq: bool = False
|
||||
|
||||
|
||||
class VisionConfig(BaseModel):
|
||||
finetune_vision_layers: bool = True
|
||||
finetune_language_layers: bool = True
|
||||
finetune_attention_modules: bool = True
|
||||
finetune_mlp_modules: bool = True
|
||||
|
||||
|
||||
class LoggingConfig(BaseModel):
|
||||
enable_wandb: bool = False
|
||||
wandb_project: str = "unsloth-training"
|
||||
wandb_token: Optional[str] = None
|
||||
enable_tensorboard: bool = False
|
||||
tensorboard_dir: str = "runs"
|
||||
|
||||
|
||||
class Config(BaseModel):
|
||||
model: Optional[str] = None
|
||||
data: DataConfig = Field(default_factory=DataConfig)
|
||||
training: TrainingConfig = Field(default_factory=TrainingConfig)
|
||||
lora: LoraConfig = Field(default_factory=LoraConfig)
|
||||
vision: VisionConfig = Field(default_factory=VisionConfig)
|
||||
logging: LoggingConfig = Field(default_factory=LoggingConfig)
|
||||
|
||||
def apply_overrides(self, **kwargs):
|
||||
"""Apply CLI overrides by matching arg names to config fields."""
|
||||
for key, value in kwargs.items():
|
||||
if value is None:
|
||||
continue
|
||||
if hasattr(self, key):
|
||||
setattr(self, key, value)
|
||||
else:
|
||||
for section in (self.data, self.training, self.lora, self.vision, self.logging):
|
||||
if hasattr(section, key):
|
||||
setattr(section, key, value)
|
||||
break
|
||||
|
||||
def model_kwargs(self, use_lora: bool, is_vision: bool) -> dict:
|
||||
"""Return kwargs for trainer.prepare_model_for_training()."""
|
||||
# Determine target modules based on model type
|
||||
if use_lora and is_vision:
|
||||
target_modules = ["all-linear"] if self.lora.vision_all_linear else []
|
||||
else:
|
||||
target_modules = [m.strip() for m in self.lora.target_modules.split(",") if m.strip()]
|
||||
|
||||
return {
|
||||
"use_lora": use_lora,
|
||||
"finetune_vision_layers": self.vision.finetune_vision_layers,
|
||||
"finetune_language_layers": self.vision.finetune_language_layers,
|
||||
"finetune_attention_modules": self.vision.finetune_attention_modules,
|
||||
"finetune_mlp_modules": self.vision.finetune_mlp_modules,
|
||||
"target_modules": target_modules,
|
||||
"lora_r": self.lora.lora_r,
|
||||
"lora_alpha": self.lora.lora_alpha,
|
||||
"lora_dropout": self.lora.lora_dropout,
|
||||
"use_gradient_checkpointing": self.training.gradient_checkpointing,
|
||||
"use_rslora": self.lora.use_rslora,
|
||||
"use_loftq": self.lora.use_loftq,
|
||||
}
|
||||
|
||||
def training_kwargs(self) -> dict:
|
||||
"""Return kwargs for trainer.start_training()."""
|
||||
return {
|
||||
"output_dir": str(self.training.output_dir),
|
||||
"num_epochs": self.training.num_epochs,
|
||||
"learning_rate": self.training.learning_rate,
|
||||
"batch_size": self.training.batch_size,
|
||||
"gradient_accumulation_steps": self.training.gradient_accumulation_steps,
|
||||
"warmup_steps": self.training.warmup_steps,
|
||||
"max_steps": self.training.max_steps,
|
||||
"save_steps": self.training.save_steps,
|
||||
"weight_decay": self.training.weight_decay,
|
||||
"random_seed": self.training.random_seed,
|
||||
"packing": self.training.packing,
|
||||
"train_on_completions": self.training.train_on_completions,
|
||||
"max_seq_length": self.training.max_seq_length,
|
||||
"enable_wandb": self.logging.enable_wandb,
|
||||
"wandb_project": self.logging.wandb_project,
|
||||
"wandb_token": self.logging.wandb_token,
|
||||
"enable_tensorboard": self.logging.enable_tensorboard,
|
||||
"tensorboard_dir": self.logging.tensorboard_dir,
|
||||
}
|
||||
|
||||
|
||||
def load_config(path: Optional[Path]) -> Config:
|
||||
"""Load config from YAML/JSON file, or return defaults if no path given."""
|
||||
if not path:
|
||||
return Config()
|
||||
|
||||
path = Path(path)
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"Config file not found: {path}")
|
||||
|
||||
text = path.read_text(encoding="utf-8")
|
||||
if path.suffix.lower() in {".yaml", ".yml"}:
|
||||
data = yaml.safe_load(text) or {}
|
||||
else:
|
||||
import json
|
||||
data = json.loads(text or "{}")
|
||||
|
||||
return Config(**data)
|
||||
Loading…
Add table
Add a link
Reference in a new issue