File size: 5,665 Bytes
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
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.")