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() def _friendly_name(value: str, fallback: str) -> str: text = str(value or "").strip() if not text: return fallback if re.search(r"^[A-Za-z]:[\\/]", text) or "/" in text or "\\" in text: return Path(text).name or fallback return text def _is_lora_training_checkpoint(path: Path) -> bool: """Return whether a LoRA weight is an intermediate training snapshot. The LoRA trainer writes both the finished adapter and periodic weights such as ``name_epoch_0050.safetensors``. The latter are useful for recovery, but are not independently selectable models in ADAM's model library. """ name = path.stem.casefold() return bool(re.search( r"(?:^|[_\- ])(?:checkpoint(?:[_\- ]?(?:epoch|e|step))?|epoch|e|step)[_\- ]?\d+(?:[_\- ]|$)", name, )) @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 = "" metadata: dict[str, Any] | None = None @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()), metadata=dict(payload.get("metadata") or {}), ) 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, metadata: dict[str, Any] | None = None, 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 metadata: current = dict(asset.metadata or {}) current.update(metadata) asset.metadata = current 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", "metadata", ) if key in item } if "trigger_word" in item: metadata = dict(values.get("metadata") or {}) metadata["trigger_word"] = str(item.get("trigger_word") or "") values["metadata"] = metadata 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, *, persist: bool = True) -> None: # Models are stored by their output folder (or the model file itself). # Keep the registry in step with the filesystem so removing an old # output cannot leave a ghost model that makes name matching ambiguous. self.assets = [ item for item in self.assets if item.kind != "model" or ( item.path.strip() and Path(item.path).expanduser().exists() ) # Old ADAM versions registered LoRA epoch snapshots. Prune those # stale records as well as skipping them during new discovery. and not ( item.trainer == "lora" and _is_lora_training_checkpoint(Path(item.path)) ) ] folders = config.get("tool_folders", {}) if not isinstance(folders, dict): return folders = dict(folders) app_root = self.path.parent.parent from adam.video_lora import discover_assets as discover_video_assets discover_video_assets(self, app_root, config) if not folders.get("oasis_trainer"): try: external = json.loads((app_root / "config" / "external_tools.json").read_text(encoding="utf-8")) except (OSError, ValueError, TypeError, json.JSONDecodeError): external = {} for entry in external.get("tools", []) if isinstance(external, dict) else []: if isinstance(entry, dict) and entry.get("id") == "external_oasis_game_trainer": root = str(entry.get("backend", {}).get("root", "")) if root: folders["oasis_trainer"] = root 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() and not _is_lora_training_checkpoint(path) ): 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"), ("oasis", "oasis_trainer", "output_action_flow_models"), ): root = Path(str(folders.get(folder_name, ""))) / output_name if not root.is_dir(): continue # LoRA Trainer versions do not all agree on their output layout. # Some write ``output//.safetensors`` while others add # a second folder below the run. Register the actual weight file # in either layout so the generator can load it directly. if trainer == "lora": for checkpoint_path in root.rglob("*.safetensors"): if ( not checkpoint_path.is_file() or "_comfy" in checkpoint_path.stem.casefold() or _is_lora_training_checkpoint(checkpoint_path) ): continue trigger_word = "" for metadata_path in (checkpoint_path.parent / "model_info.json", checkpoint_path.parent.parent / "model_info.json"): try: metadata = json.loads(metadata_path.read_text(encoding="utf-8")) trigger_word = str(metadata.get("trigger_word") or "") if trigger_word: break except (OSError, ValueError, TypeError, json.JSONDecodeError): continue self.register( kind="model", name=checkpoint_path.stem.removesuffix("_cancelled"), path=str(checkpoint_path), trainer="lora", checkpoint=str(checkpoint_path), metadata={"trigger_word": trigger_word or checkpoint_path.stem.removesuffix("_cancelled")}, persist=False, ) continue for folder in root.iterdir(): if not folder.is_dir(): continue name = folder.name dataset_path = "" if trainer == "ddpm": # DDPM writes a durable sidecar with the friendly model name and # source dataset. Prefer it over a filesystem-safe folder name. 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 == "flow": 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 else: checkpoints = [] try: metadata = json.loads( (folder / "action_flow_model_info.json").read_text(encoding="utf-8") ) if metadata.get("model_type") != "action_conditioned_rectified_flow_video": continue if not (folder / "unet" / "config.json").is_file(): continue name = _friendly_name( str(metadata.get("model_name") or metadata.get("name") or name), folder.name, ) except (OSError, ValueError, TypeError, json.JSONDecodeError): continue checkpoint = ( str(folder) if trainer in {"flow", "oasis"} 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, metadata=({"trigger_word": trigger_word or name} if trainer == "lora" else None), persist=False, ) try: from adam.dataset_registry import DatasetRegistry dataset_registry = DatasetRegistry(app_root, config) valid_location_ids = {location.id for location in dataset_registry.known_locations()} self.assets = [ item for item in self.assets if not ( item.kind == "dataset" and isinstance(item.metadata, dict) and item.metadata.get("dataset_registry_source") in {"adam", "tool"} and item.metadata.get("dataset_location_id") and item.metadata.get("dataset_location_id") not in valid_location_ids ) ] dataset_registry.discover_into_assets(self, persist=False, update_cache=persist) except Exception: pass if persist: 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