SyntheticMDProductions's picture
Some of Adams structure
e0265b9 verified
Raw
History Blame Contribute Delete
5.43 kB
from __future__ import annotations
import re
from dataclasses import dataclass
from pathlib import Path
from adam.config import ConfigManager
@dataclass(frozen=True, slots=True)
class ToolFolderDefinition:
tool_id: str
name: str
expected_files: tuple[str, ...] = ()
@dataclass(frozen=True, slots=True)
class ToolFolderStatus:
tool_id: str
name: str
path: str
exists: bool
valid: bool
entry_points: tuple[str, ...]
message: str
TOOL_FOLDER_DEFINITIONS = (
ToolFolderDefinition(
"dataset_collector",
"Dataset Collector",
("collector.py", "dataset_collector.py", "main.py"),
),
ToolFolderDefinition(
"caption_generator",
"Caption Generator",
("caption.py", "caption_generator.py", "main.py"),
),
ToolFolderDefinition(
"lora_trainer",
"LoRA Trainer",
("train.py", "lora_train.py", "main.py", "src/loratrainer/main.py"),
),
ToolFolderDefinition(
"ddpm_trainer",
"DDPM Trainer",
("train.py", "appStableDiffusion.py"),
),
ToolFolderDefinition(
"flow_trainer",
"Flow Matching Trainer",
("flow_matching_app.py", "roblox_action_flow_app.py"),
),
ToolFolderDefinition(
"preview_generator",
"Preview Generator",
("generate.py", "preview.py", "main.py"),
),
)
class ToolFolderManager:
def __init__(self, config: ConfigManager) -> None:
self.config = config
self.definitions = {
definition.tool_id: definition
for definition in TOOL_FOLDER_DEFINITIONS
}
def paths(self) -> dict[str, str]:
stored = self.config.get("tool_folders", {})
return dict(stored) if isinstance(stored, dict) else {}
def get(self, tool_id: str) -> str:
return str(self.paths().get(tool_id, ""))
def set(self, tool_id: str, path: str) -> ToolFolderStatus:
if tool_id not in self.definitions:
raise KeyError(f"Unknown tool folder: {tool_id}")
stored = self.paths()
stored[tool_id] = path.strip().strip('"')
self.config.update({"tool_folders": stored})
return self.scan(tool_id)
def update(self, values: dict[str, str]) -> dict[str, ToolFolderStatus]:
stored = self.paths()
for tool_id, value in values.items():
if tool_id in self.definitions:
stored[tool_id] = value.strip().strip('"')
self.config.update({"tool_folders": stored})
return {tool_id: self.scan(tool_id) for tool_id in values}
def scan(self, tool_id: str) -> ToolFolderStatus:
definition = self.definitions[tool_id]
raw_path = self.get(tool_id)
if not raw_path:
return ToolFolderStatus(
tool_id,
definition.name,
"",
False,
False,
(),
"Not configured",
)
folder = Path(raw_path).expanduser()
if not folder.is_dir():
return ToolFolderStatus(
tool_id,
definition.name,
str(folder),
False,
False,
(),
"Folder not found",
)
entry_points = tuple(
filename
for filename in definition.expected_files
if (folder / filename).is_file()
)
if entry_points:
message = f"Detected 路 {', '.join(entry_points)}"
valid = True
else:
top_level_python = sorted(path.name for path in folder.glob("*.py"))
entry_points = tuple(top_level_python[:5])
valid = bool(entry_points)
message = (
f"Python project detected 路 review {entry_points[0]}"
if entry_points
else "Folder found 路 no Python entry point detected"
)
return ToolFolderStatus(
tool_id,
definition.name,
str(folder.resolve()),
True,
valid,
entry_points,
message,
)
def scan_all(self) -> dict[str, ToolFolderStatus]:
return {
tool_id: self.scan(tool_id)
for tool_id in self.definitions
}
def parse_assignments(self, text: str) -> dict[str, str]:
"""Recognize folder assignments pasted into chat without executing them."""
patterns = {
"ddpm_trainer": r"(?im)^\s*DDPM(?:\s+Trainer)?\s*:\s*(.+?)\s*$",
"flow_trainer": (
r"(?im)^\s*Flow(?:\s+Matching)?(?:\s+Trainer)?\s*:\s*(.+?)\s*$"
),
"lora_trainer": r"(?im)^\s*LoRA(?:\s+Trainer)?\s*:\s*(.+?)\s*$",
"dataset_collector": (
r"(?im)^\s*Dataset(?:\s+Collector)?\s*:\s*(.+?)\s*$"
),
"caption_generator": (
r"(?im)^\s*Caption(?:\s+Generator)?\s*:\s*(.+?)\s*$"
),
"preview_generator": (
r"(?im)^\s*Preview(?:\s+Generator)?\s*:\s*(.+?)\s*$"
),
}
assignments: dict[str, str] = {}
for tool_id, pattern in patterns.items():
match = re.search(pattern, text)
if match:
assignments[tool_id] = match.group(1).strip().strip('"')
return assignments