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