from __future__ import annotations from dataclasses import dataclass from typing import Any 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 = "" 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", "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(), 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.") if command.trainer not in {"ddpm", "lora", "flow"}: raise CommandValidationError("Trainer must be ddpm, flow, or lora.") 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.") command._validate_options() return command def _validate_options(self) -> None: options = self.training_options or {} 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", }, "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"}, }[self.trainer] unknown = set(options) - allowed if unknown: raise CommandValidationError(f"Unsupported {self.trainer} training options: {', '.join(sorted(unknown))}") 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.")