SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
5.67 kB
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.")