SyntheticMDProductions's picture
Update ADAM safety, UI, and model workflows (#1)
c61c435
Raw History Blame Contribute Delete
1.05 kB
from __future__ import annotations
from pathlib import Path
from typing import Any
from .generic import GenericModelInspector
class MaskGITInspector(GenericModelInspector):
architecture = "MaskGIT"
def inspect(self, path: str | Path, *, recorded_architecture: str = "", run_settings: dict[str, Any] | None = None, progress=None, cancelled=None):
summary = super().inspect(
path,
recorded_architecture=recorded_architecture or "maskgit",
run_settings=run_settings,
progress=progress,
cancelled=cancelled,
)
summary.architecture = "MaskGIT" if summary.architecture == "Generic / Unknown" else summary.architecture
names = "\n".join(tensor.name for tensor in summary.tensors).casefold()
if summary.tensors and not any(token in names for token in ("attention", "attn", "transformer", "embed")):
summary.health.append("Worth inspecting: expected MaskGIT transformer or embedding tensors were not obvious")
return summary