SyntheticMDProductions's picture
ADAM October 2026 source release: PixelRow, INRFlow, Wan Video, Oasis player and field guide
f8c73f9 verified
Raw History Blame Contribute Delete
6.97 kB
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from adam.model_plugins import ModelPluginRegistry, validate_settings
class CommandValidationError(ValueError):
pass
@dataclass(frozen=True, slots=True)
class TrainingCommand:
action: str
trainer: str
dataset: str
model_name: str
epochs: int
output: str = "default output"
resume_from: str = ""
base_model: str = ""
trigger_word: str = ""
training_options: dict[str, Any] | None = None
@classmethod
def from_dict(cls, payload: dict[str, Any]) -> "TrainingCommand":
allowed = {
"action", "trainer", "dataset", "model_name", "epochs", "output",
"resume_from", "base_model", "trigger_word",
"training_options",
}
unknown = set(payload) - allowed
if unknown:
raise CommandValidationError(
f"Unsupported command fields: {', '.join(sorted(unknown))}"
)
try:
command = cls(
action=str(payload.get("action", "train")).casefold(),
trainer=str(payload.get("trainer", "")).casefold(),
dataset=str(payload.get("dataset", "")).strip(),
model_name=str(payload.get("model_name", "")).strip(),
epochs=int(payload.get("epochs", 0) or 0),
output=str(payload.get("output", "default output")).strip(),
resume_from=str(payload.get("resume_from", "")).strip(),
base_model=str(payload.get("base_model", "")).strip(),
trigger_word=str(
payload.get("trigger_word")
or (payload.get("training_options") or {}).get("trigger_word")
or ""
).strip(),
training_options=dict(payload.get("training_options") or {}),
)
except (TypeError, ValueError) as exc:
raise CommandValidationError("Training command fields have invalid types.") from exc
if command.action not in {"train", "resume_training"}:
raise CommandValidationError("Training action must be train or resume_training.")
plugin_schema = ModelPluginRegistry(Path.cwd()).training_schema(command.trainer)
if command.trainer not in {"ddpm", "lora", "flow"} and not plugin_schema:
raise CommandValidationError("Trainer must be a discovered model plugin.")
if not command.dataset or not command.model_name:
raise CommandValidationError("Dataset and model name are required.")
if not 1 <= command.epochs <= 100_000:
raise CommandValidationError("Epoch count must be between 1 and 100000.")
if command.action == "resume_training" and not command.resume_from:
raise CommandValidationError("Resume training requires an explicit checkpoint.")
if command.trainer == "lora":
trigger = command.trigger_word or command.model_name
if len(trigger) > 128 or any(char in trigger for char in '<>:"/\\|?*\x00'):
raise CommandValidationError("LoRA trigger word must be short text without reserved characters.")
command._validate_options()
return command
def _validate_options(self) -> None:
options = self.training_options or {}
schema = ModelPluginRegistry(Path.cwd()).training_schema(self.trainer)
allowed = set(schema)
if not allowed:
allowed = {
"ddpm": {
"resolution", "batch_size", "learning_rate", "gradient_accumulation_steps",
"dataloader_num_workers", "mixed_precision", "save_every", "preview_steps",
"training_intensity", "preview_enabled", "preview_every", "preview_prompt",
"preview_seed", "training_aspect_ratio", "resize_mode",
},
"flow": {
"resolution", "batch_size", "learning_rate", "gradient_accumulation",
"workers", "mixed_precision", "save_every", "preview_every", "preview_steps",
"gradient_checkpointing", "preview_enabled", "preview_prompt", "preview_seed",
},
"lora": {"preview_enabled", "preview_every", "preview_prompt", "preview_seed", "trigger_word"},
}[self.trainer]
unknown = set(options) - allowed
if unknown:
raise CommandValidationError(f"Unsupported {self.trainer} training options: {', '.join(sorted(unknown))}")
if schema:
errors = validate_settings(
{key: spec for key, spec in schema.items() if key in options},
options,
)
if errors:
raise CommandValidationError(" ".join(errors))
return
integer_ranges = {
"resolution": (64, 512), "batch_size": (1, 64),
"gradient_accumulation_steps": (1, 64), "gradient_accumulation": (1, 64),
"dataloader_num_workers": (0, 16), "workers": (0, 16),
"save_every": (1, 1000), "preview_every": (1, 100_000),
"preview_steps": (1, 500), "training_intensity": (10, 100),
}
for key, (low, high) in integer_ranges.items():
if key in options and (not isinstance(options[key], int) or not low <= options[key] <= high):
raise CommandValidationError(f"{key} must be an integer between {low} and {high}.")
if "learning_rate" in options:
value = options["learning_rate"]
if not isinstance(value, (int, float)) or isinstance(value, bool) or not 1e-7 <= float(value) <= 0.1:
raise CommandValidationError("learning_rate must be between 0.0000001 and 0.1.")
if "mixed_precision" in options and options["mixed_precision"] not in {"fp16", "no"}:
raise CommandValidationError("mixed_precision must be fp16 or no.")
if "gradient_checkpointing" in options and not isinstance(options["gradient_checkpointing"], bool):
raise CommandValidationError("gradient_checkpointing must be true or false.")
if "preview_enabled" in options and not isinstance(options["preview_enabled"], bool):
raise CommandValidationError("preview_enabled must be true or false.")
if "preview_seed" in options and (
not isinstance(options["preview_seed"], int)
or isinstance(options["preview_seed"], bool)
or not 0 <= options["preview_seed"] <= 2_147_483_647
):
raise CommandValidationError("preview_seed must be an integer between 0 and 2147483647.")
if "preview_prompt" in options and (
not isinstance(options["preview_prompt"], str) or len(options["preview_prompt"]) > 2_000
):
raise CommandValidationError("preview_prompt must be text up to 2000 characters.")