| from __future__ import annotations |
|
|
| import json |
| import re |
| from dataclasses import asdict, dataclass |
| from datetime import datetime, timezone |
| from pathlib import Path |
| from typing import Any |
| from uuid import uuid4 |
|
|
|
|
| def _now() -> str: |
| return datetime.now(timezone.utc).isoformat() |
|
|
|
|
| def _normal(value: str) -> str: |
| return re.sub(r"[^a-z0-9]+", " ", value.casefold()).strip() |
|
|
|
|
| @dataclass(slots=True) |
| class Asset: |
| id: str |
| kind: str |
| name: str |
| path: str |
| trainer: str = "" |
| dataset_id: str = "" |
| checkpoint: str = "" |
| epochs: int = 0 |
| created_at: str = "" |
|
|
| @classmethod |
| def from_dict(cls, payload: dict[str, Any]) -> "Asset": |
| return cls( |
| id=str(payload.get("id") or uuid4().hex[:12]), |
| kind=str(payload.get("kind", "")), |
| name=str(payload.get("name", "")), |
| path=str(payload.get("path", "")), |
| trainer=str(payload.get("trainer", "")), |
| dataset_id=str(payload.get("dataset_id", "")), |
| checkpoint=str(payload.get("checkpoint", "")), |
| epochs=int(payload.get("epochs", 0) or 0), |
| created_at=str(payload.get("created_at") or _now()), |
| ) |
|
|
|
|
| class AssetRegistry: |
| """Persistent friendly-name index for datasets, models, and checkpoints.""" |
|
|
| def __init__(self, root: Path) -> None: |
| self.path = root.resolve() / "data" / "assets.json" |
| self.assets: list[Asset] = [] |
| self.load() |
|
|
| def load(self) -> None: |
| try: |
| payload = json.loads(self.path.read_text(encoding="utf-8")) |
| self.assets = [ |
| Asset.from_dict(item) |
| for item in payload.get("assets", []) |
| if isinstance(item, dict) |
| ] |
| except (OSError, ValueError, TypeError, json.JSONDecodeError): |
| self.assets = [] |
|
|
| def save(self) -> None: |
| self.path.parent.mkdir(parents=True, exist_ok=True) |
| temporary = self.path.with_suffix(".tmp") |
| temporary.write_text( |
| json.dumps({"assets": [asdict(item) for item in self.assets]}, indent=2), |
| encoding="utf-8", |
| ) |
| temporary.replace(self.path) |
|
|
| def register( |
| self, |
| *, |
| kind: str, |
| name: str, |
| path: str, |
| trainer: str = "", |
| dataset_id: str = "", |
| checkpoint: str = "", |
| epochs: int = 0, |
| persist: bool = True, |
| ) -> Asset: |
| resolved = str(Path(path).expanduser().resolve()) |
| existing = next( |
| ( |
| item |
| for item in self.assets |
| if item.kind == kind and Path(item.path) == Path(resolved) |
| ), |
| None, |
| ) |
| asset = existing or Asset(uuid4().hex[:12], kind, name, resolved) |
| asset.name = name.strip() or Path(resolved).name |
| asset.trainer = trainer |
| asset.dataset_id = dataset_id |
| asset.checkpoint = checkpoint |
| asset.epochs = int(epochs) |
| asset.created_at = asset.created_at or _now() |
| if existing is None: |
| self.assets.insert(0, asset) |
| if persist: |
| self.save() |
| return asset |
|
|
| def ingest_result(self, result: dict[str, Any]) -> None: |
| entries = result.get("assets", []) |
| if not isinstance(entries, list): |
| return |
| for item in entries: |
| if not isinstance(item, dict): |
| continue |
| if item.get("kind") and item.get("path"): |
| values = { |
| key: item[key] |
| for key in ( |
| "kind", "name", "path", "trainer", "dataset_id", |
| "checkpoint", "epochs", |
| ) |
| if key in item |
| } |
| dataset_path = str(item.get("dataset_path", "")) |
| if item.get("kind") == "model" and dataset_path and Path(dataset_path).is_dir(): |
| dataset = self.register( |
| kind="dataset", |
| name=Path(dataset_path).name, |
| path=dataset_path, |
| ) |
| values["dataset_id"] = dataset.id |
| values.setdefault("name", Path(str(item["path"])).name) |
| self.register(**values) |
|
|
| def find(self, kind: str, query: str, *, trainer: str = "") -> list[Asset]: |
| wanted = _normal(query) |
| matches = [] |
| exact = [] |
| for item in self.assets: |
| if item.kind != kind or (trainer and item.trainer != trainer): |
| continue |
| haystacks = {_normal(item.name), _normal(Path(item.path).name)} |
| if wanted in haystacks: |
| exact.append(item) |
| elif any(wanted and wanted in value for value in haystacks): |
| matches.append(item) |
| return exact or matches |
|
|
| def discover(self, config: Any) -> None: |
| folders = config.get("tool_folders", {}) |
| if not isinstance(folders, dict): |
| return |
| app_root = self.path.parent.parent |
| external_lora_root = app_root / "LoRAModelsHere" |
| if external_lora_root.is_dir(): |
| for path in external_lora_root.rglob("*.safetensors"): |
| if path.is_file() and "_comfy" not in path.stem.casefold(): |
| self.register( |
| kind="model", |
| name=path.stem.removesuffix("_cancelled"), |
| path=str(path), |
| trainer="lora", |
| checkpoint=str(path), |
| persist=False, |
| ) |
| base_model_root = app_root / "LoRA StableDiffusionModels Here" |
| if base_model_root.is_dir(): |
| for path in base_model_root.iterdir(): |
| is_model_file = path.is_file() and path.suffix.casefold() in { |
| ".safetensors", ".ckpt", ".pt", ".bin" |
| } |
| is_diffusers_folder = path.is_dir() and ( |
| (path / "model_index.json").is_file() |
| or (path / "unet" / "config.json").is_file() |
| ) |
| if is_model_file or is_diffusers_folder: |
| self.register( |
| kind="base_model", |
| name=path.stem if path.is_file() else path.name, |
| path=str(path), |
| trainer="stable_diffusion", |
| persist=False, |
| ) |
| flow_datasets = self._flow_dataset_paths() |
| collector = Path(str(folders.get("dataset_collector", ""))) / "Datasets" |
| if collector.is_dir(): |
| for folder in collector.iterdir(): |
| if folder.is_dir(): |
| self.register( |
| kind="dataset", name=folder.name, path=str(folder), persist=False |
| ) |
| for trainer, folder_name, output_name in ( |
| ("ddpm", "ddpm_trainer", "output"), |
| ("lora", "lora_trainer", "output"), |
| ("flow", "flow_trainer", "output_flow_models"), |
| ): |
| root = Path(str(folders.get(folder_name, ""))) / output_name |
| if not root.is_dir(): |
| continue |
| for folder in root.iterdir(): |
| if not folder.is_dir(): |
| continue |
| name = folder.name |
| dataset_path = "" |
| if trainer == "ddpm": |
| |
| |
| try: |
| metadata = json.loads((folder / "model_info.json").read_text(encoding="utf-8")) |
| name = str(metadata.get("model_name") or metadata.get("name") or name) |
| dataset_path = str(metadata.get("dataset_dir") or "") |
| except (OSError, ValueError, TypeError, json.JSONDecodeError): |
| pass |
| checkpoints = sorted( |
| folder.glob("checkpoint-*"), |
| key=lambda p: int(p.name.rsplit("-", 1)[-1]) |
| if p.name.rsplit("-", 1)[-1].isdigit() |
| else -1, |
| ) |
| elif trainer == "lora": |
| checkpoints = sorted( |
| ( |
| path for path in folder.glob("*.safetensors") |
| if "_comfy" not in path.stem.casefold() |
| ), |
| key=lambda p: p.stat().st_mtime, |
| ) |
| if checkpoints: |
| name = checkpoints[-1].stem.removesuffix("_cancelled") |
| else: |
| checkpoints = [] |
| try: |
| metadata = json.loads( |
| (folder / "flow_model_info.json").read_text(encoding="utf-8") |
| ) |
| if metadata.get("model_type") != "rectified_flow": |
| continue |
| if not (folder / "unet" / "config.json").is_file(): |
| continue |
| name = str(metadata.get("model_name") or metadata.get("name") or name) |
| dataset_path = flow_datasets.get(str(folder.resolve()), "") |
| except (OSError, ValueError, TypeError, json.JSONDecodeError): |
| continue |
| checkpoint = ( |
| str(folder) if trainer == "flow" else str(checkpoints[-1]) if checkpoints else "" |
| ) |
| dataset_id = "" |
| if dataset_path and Path(dataset_path).is_dir(): |
| dataset_asset = self.register( |
| kind="dataset", |
| name=Path(dataset_path).name, |
| path=dataset_path, |
| persist=False, |
| ) |
| dataset_id = dataset_asset.id |
| self.register( |
| kind="model", |
| name=name, |
| path=str(folder), |
| trainer=trainer, |
| dataset_id=dataset_id, |
| checkpoint=checkpoint, |
| persist=False, |
| ) |
| self.save() |
|
|
| def _flow_dataset_paths(self) -> dict[str, str]: |
| """Recover source datasets for Flow models created by ADAM in older runs.""" |
| jobs_path = self.path.parent / "jobs.json" |
| try: |
| payload = json.loads(jobs_path.read_text(encoding="utf-8")) |
| jobs = payload.get("jobs", []) |
| except (OSError, ValueError, TypeError, json.JSONDecodeError): |
| return {} |
| links: dict[str, str] = {} |
| if not isinstance(jobs, list): |
| return links |
| for job in jobs: |
| if not isinstance(job, dict): |
| continue |
| plan = job.get("plan", {}) |
| steps = plan.get("steps", []) if isinstance(plan, dict) else [] |
| if not isinstance(steps, list): |
| continue |
| for step in steps: |
| if not isinstance(step, dict) or step.get("tool_id") != "flow_trainer": |
| continue |
| arguments = step.get("arguments", {}) |
| if not isinstance(arguments, dict): |
| continue |
| output = str(arguments.get("output_dir", "")) |
| dataset = str(arguments.get("dataset_dir", "")) |
| if not output or not dataset or not Path(dataset).is_dir(): |
| continue |
| try: |
| links[str(Path(output).expanduser().resolve())] = str(Path(dataset).expanduser().resolve()) |
| except OSError: |
| continue |
| return links |
|
|