AI_Development_Automation_Manager / adam /training_assistant.py
SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
16.1 kB
from __future__ import annotations
import json
import re
import shutil
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from adam.models import ExecutionPlan
from adam.orion import apply_orion_review
DEFAULT_PRESETS: dict[str, dict[str, Any]] = {
"Character LoRA": {
"trainer": "lora", "epochs": 100, "image_count": 60,
"description": "A balanced starting point for a character or person.",
},
"Style LoRA": {
"trainer": "lora", "epochs": 80, "image_count": 80,
"description": "A broader image set for learning a visual style.",
},
"DDPM Test Run": {
"trainer": "ddpm", "epochs": 25, "image_count": 40,
"description": "A short run to verify the dataset and training setup.",
},
"DDPM Full Run": {
"trainer": "ddpm", "epochs": 100, "image_count": 100,
"description": "A practical default for a full DDPM experiment.",
},
"Flow Test Run": {
"trainer": "flow", "epochs": 25, "image_count": 40,
"description": "A short Flow Matching setup check.",
},
}
def parse_model_batch_names(text: str) -> list[str]:
"""Return unique, user-ordered model subjects from a pasted line list."""
names: list[str] = []
seen: set[str] = set()
for raw in text.splitlines():
name = re.sub(r"^\s*(?:[-*•]|\d+[.)])\s*", "", raw).strip()
key = re.sub(r"\s+", " ", name).casefold()
if name and key not in seen:
names.append(re.sub(r"\s+", " ", name))
seen.add(key)
return names
def build_dataset_collection_request(
subject: str,
*,
image_count: int = 100,
collection_mode: str = "target",
) -> str:
"""Build the dataset-only first phase used by a saved model batch."""
subject = subject.strip()
if not subject:
raise ValueError("Dataset collection requires a subject.")
if collection_mode == "all_available":
return (
f"Collect a dataset of {subject} with as many available images as Bing "
"returns (up to 5,000)."
)
return f"Collect a dataset of {image_count} images of {subject}."
@dataclass(slots=True)
class PreflightItem:
level: str
message: str
@dataclass(slots=True)
class DatasetMatch:
status: str
dataset_name: str = ""
score: float = 0.0
_DATASET_NAME_NOISE = {
"dataset", "datasets", "image", "images", "picture", "pictures",
"photo", "photos", "collection", "collected",
}
def _dataset_name_tokens(value: object) -> set[str]:
words = re.findall(r"[a-z0-9]+", str(value).casefold())
return {
word[:-1] if word.endswith("s") and len(word) > 3 else word
for word in words if word not in _DATASET_NAME_NOISE
}
def suggest_existing_dataset(state: dict[str, object], datasets: list[Any]) -> DatasetMatch:
"""Safely match one batch model to a registered dataset by its human name."""
queries = [
_dataset_name_tokens(state.get("model_name", "")),
_dataset_name_tokens(state.get("subject", "")),
]
queries = [query for query in queries if query]
if not queries:
return DatasetMatch("unmatched")
scored: list[tuple[float, Any]] = []
for asset in datasets:
name = str(getattr(asset, "name", ""))
path = Path(str(getattr(asset, "path", "")))
tokens = _dataset_name_tokens(name)
if not name or not tokens or not path.is_dir():
continue
score = 0.0
for query in queries:
overlap = len(query & tokens) / len(query)
extra_penalty = min(0.20, len(tokens - query) * 0.08)
score = max(score, overlap - extra_penalty)
if score >= 0.80:
scored.append((score, asset))
if not scored:
return DatasetMatch("unmatched")
scored.sort(key=lambda item: (-item[0], len(str(getattr(item[1], "name", "")))))
best_score, best = scored[0]
if len(scored) > 1 and best_score - scored[1][0] < 0.10:
return DatasetMatch("ambiguous", score=best_score)
return DatasetMatch("matched", str(getattr(best, "name", "")), best_score)
def combine_training_plans(plans: list[ExecutionPlan]) -> ExecutionPlan:
"""Combine independently validated model plans into one sequential job."""
usable = [plan for plan in plans if plan.steps]
if not usable:
raise ValueError("A training batch needs at least one actionable model plan.")
if len(usable) == 1:
return usable[0]
summaries = [
f"{index}. {plan.project_name}: {plan.summary.splitlines()[0]}"
for index, plan in enumerate(usable, 1)
]
reasons = [plan.confirmation_reason for plan in usable if plan.confirmation_reason]
has_training = any(
step.tool_id.endswith("_trainer")
for plan in usable for step in plan.steps
)
batch_kind = "training" if has_training else "dataset collection"
return ExecutionPlan(
request="\n\n".join(plan.request for plan in usable),
summary=(
f"Sequential {batch_kind} batch with {len(usable)} items. ADAM will finish "
"each item before starting the next; a failed step stops the batch.\n\n"
+ "\n".join(summaries)
),
steps=[step for plan in usable for step in plan.steps],
requires_confirmation=any(plan.requires_confirmation for plan in usable),
confirmation_reason="; ".join(dict.fromkeys(reasons)) or (
"This batch contains multiple model workflows. Review every model and its "
"output path before starting."
),
project_name=(
f"Training batch ({len(usable)} models)" if has_training
else f"Dataset collection batch ({len(usable)} datasets)"
),
)
def estimate_plan(plan: Any) -> list[PreflightItem]:
"""Add deliberately conservative, clearly labelled planning estimates."""
estimates: list[PreflightItem] = []
for step in plan.steps:
if not step.tool_id.endswith("_trainer"):
continue
trainer = step.tool_id.removesuffix("_trainer")
epochs = max(1, int(step.arguments.get("epochs", 1) or 1))
dataset = Path(str(step.arguments.get("dataset_dir", ""))).expanduser()
image_count = 0
if dataset.is_dir():
try:
image_count = sum(
1
for path in dataset.rglob("*")
if path.is_file()
and path.suffix.casefold()
in {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
)
except OSError:
image_count = 0
image_count = image_count or 60
workload = epochs * image_count
seconds_per_image_epoch = {
"lora": 0.12,
"ddpm": 0.07,
"flow": 0.10,
}.get(trainer, 0.10)
center_minutes = max(1, int(workload * seconds_per_image_epoch / 60))
low = max(1, center_minutes // 2)
high = max(low + 1, center_minutes * 3)
checkpoint_gb = {
"lora": 0.25,
"ddpm": 1.0,
"flow": 1.0,
}.get(trainer, 0.75)
checkpoint_count = max(1, min(20, epochs // 25 + 1))
disk_gb = checkpoint_gb * checkpoint_count
typical_vram = {"lora": 8, "ddpm": 6, "flow": 8}.get(trainer, 8)
estimates.extend(
[
PreflightItem(
"estimate",
f"Estimated workload: {workload:,} image-epochs "
f"({epochs:,} epochs × about {image_count:,} images)",
),
PreflightItem(
"estimate",
f"Rough duration: {low}{high} minutes; model size, resolution, "
"batch size, and GPU can change this substantially",
),
PreflightItem(
"estimate",
f"Suggested capacity: about {typical_vram} GB VRAM and "
f"{disk_gb:.1f} GB free for checkpoints",
),
]
)
return estimates
def presets_from_config(config: Any) -> dict[str, dict[str, Any]]:
presets = {name: dict(values) for name, values in DEFAULT_PRESETS.items()}
stored = config.get("training_presets", {})
if isinstance(stored, dict):
for name, values in stored.items():
if isinstance(name, str) and isinstance(values, dict):
presets[name] = dict(values)
return presets
def build_training_request(
*,
trainer: str,
subject: str,
dataset_name: str,
create_dataset: bool,
epochs: int,
image_count: int,
model_name: str,
collection_mode: str = "target",
training_options: dict[str, Any] | None = None,
) -> str:
subject = subject.strip()
dataset_name = dataset_name.strip()
model_name = model_name.strip() or subject or dataset_name
trainer_label = {"lora": "LoRA", "ddpm": "DDPM", "flow": "Flow Matching"}[trainer]
if create_dataset:
collection_phrase = (
"as many available images as Bing returns (up to 5,000)"
if collection_mode == "all_available"
else f"up to {image_count} images"
)
if trainer == "lora":
request = (
f"Create and train a LoRA of {subject} for {epochs} epochs "
f"using {collection_phrase}. Name the model {model_name}."
)
elif trainer == "ddpm":
request = (
f"Grab a dataset of {subject} off the internet with {collection_phrase}, "
f"name the model {model_name}, train it on a DDPM for {epochs} epochs, "
"and save it to the DDPM output."
)
else:
request = (
f"Collect a dataset of {collection_phrase} of {subject}. Then train the "
f"{subject} dataset with Flow Matching for {epochs} epochs and name the model {model_name}."
)
else:
request = (
f"From the {dataset_name} dataset, train a {trainer_label} model for {epochs} epochs. "
f"Name the model {model_name}."
)
if training_options:
request += " [ADAM_TRAINING_OPTIONS:" + json.dumps(training_options, sort_keys=True) + "]"
return request
def build_fine_tune_request(
*,
model_name: str,
trainer: str,
epochs: int,
dataset_mode: str = "original",
dataset_name: str = "",
new_subject: str = "",
image_count: int = 60,
training_options: dict[str, Any] | None = None,
) -> str:
"""Build the explicit continuation request used by the Fine-Tune assistant."""
labels = {"lora": "LoRA", "ddpm": "DDPM", "flow": "Flow Matching"}
if trainer not in labels:
raise ValueError("Fine-tuning requires a supported trainer.")
if not model_name.strip():
raise ValueError("Fine-tuning requires a model name.")
if epochs < 1:
raise ValueError("Fine-tuning requires at least one additional epoch.")
if dataset_mode not in {"original", "existing", "new"}:
raise ValueError("Fine-tuning requires a valid dataset choice.")
payload = {
"model_name": model_name.strip(),
"trainer": trainer,
"epochs": epochs,
"dataset_mode": dataset_mode,
"dataset_name": dataset_name.strip(),
"new_subject": new_subject.strip(),
"image_count": max(10, min(int(image_count), 5000)),
"training_options": dict(training_options or {}),
}
return (
f"Fine-tune {model_name.strip()} for {epochs} epochs with {labels[trainer]}. "
"[ADAM_FINE_TUNE:" + json.dumps(payload, sort_keys=True) + "]"
)
def inspect_plan(plan: Any, config: Any) -> list[PreflightItem]:
items: list[PreflightItem] = []
folders = config.get("tool_folders", {})
folders = folders if isinstance(folders, dict) else {}
checked_tools: set[str] = set()
checked_paths: set[str] = set()
for step in plan.steps:
if step.tool_id.endswith("_trainer") or step.tool_id == "dataset_collector":
if step.tool_id not in checked_tools:
configured = Path(str(folders.get(step.tool_id, ""))).expanduser()
if configured.is_dir():
items.append(PreflightItem("ready", f"{step.title}: connected"))
else:
items.append(PreflightItem("warning", f"{step.title}: program folder is not connected"))
checked_tools.add(step.tool_id)
dataset = str(step.arguments.get("dataset_dir", ""))
if dataset and dataset not in checked_paths:
if Path(dataset).is_dir():
image_count = sum(
1 for path in Path(dataset).iterdir()
if path.suffix.casefold() in {".jpg", ".jpeg", ".png", ".webp", ".bmp"}
)
detail = f"{image_count} images found" if image_count else "folder found; no top-level images detected"
items.append(PreflightItem("ready" if image_count else "warning", f"Dataset: {detail}"))
elif not any(
prior.tool_id == "dataset_collector"
and prior.arguments.get("output_dir") == dataset
for prior in plan.steps
):
items.append(PreflightItem("warning", "Dataset folder does not exist yet"))
checked_paths.add(dataset)
base_model = str(step.arguments.get("base_model", ""))
if step.tool_id == "lora_trainer":
items.append(
PreflightItem(
"ready" if base_model and Path(base_model).is_file() else "warning",
"LoRA base model is available" if base_model and Path(base_model).is_file()
else "LoRA base model still needs to be selected",
)
)
output = str(step.arguments.get("output_dir", ""))
if output:
probe = Path(output)
while not probe.exists() and probe.parent != probe:
probe = probe.parent
try:
free_gb = shutil.disk_usage(probe).free / (1024 ** 3)
items.append(
PreflightItem(
"ready" if free_gb >= 10 else "warning",
f"Output drive has {free_gb:.1f} GB free",
)
)
except OSError:
items.append(PreflightItem("warning", "Output drive space could not be checked"))
return items
def append_preflight_summary(plan: Any, config: Any) -> None:
if not plan.steps:
return
if "Pre-flight:" not in plan.summary:
items = inspect_plan(plan, config) + estimate_plan(plan)
if items:
lines = [
(
"Ready"
if item.level == "ready"
else "Estimate"
if item.level == "estimate"
else "Check"
)
+ f": {item.message}"
for item in items
]
plan.summary += "\n\nPre-flight:\n" + "\n".join(f"• {line}" for line in lines)
if not getattr(plan, "orion_review", None) and "ORION —" not in plan.summary:
apply_orion_review(plan)
def completion_recommendation(plan: Any) -> str:
tools = {step.tool_id for step in plan.steps}
if "lora_trainer" in tools or "ddpm_trainer" in tools or "flow_trainer" in tools:
return (
"Recommended next step: generate a few preview images and compare them with "
"the training dataset. If the subject is weak, improve the dataset before adding epochs."
)
if "dataset_collector" in tools:
return (
"Recommended next step: review the images and captions, remove weak or duplicate "
"examples, then open the Model Creation Assistant to start a short test run."
)
return ""