| 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.") |
|
|