File size: 6,971 Bytes
e0265b9 c61c435 e0265b9 c61c435 e0265b9 c61c435 e0265b9 c61c435 e0265b9 c61c435 e0265b9 c61c435 e0265b9 c61c435 e0265b9 c61c435 f8c73f9 c61c435 e0265b9 c61c435 e0265b9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 | 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.")
|