Spaces:
Paused
Paused
| """ | |
| models/__init__.py | |
| ────────────────── | |
| Factory that maps the loader name string (from ModelConfig.loader) | |
| to the actual class, so pipeline.py never needs to import concrete loaders. | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| from models.base_loader import BaseLoader | |
| from models.background_removal import BiRefNetLoader | |
| from models.depth_estimation import TransformersDepthLoader | |
| from models.reconstruction import ( | |
| Open3DReconstructionLoader, | |
| GaussianSplatLoader, | |
| DepthSplatLoader, | |
| ) | |
| _LOADER_MAP: dict[str, type[BaseLoader]] = { | |
| "BiRefNetLoader": BiRefNetLoader, | |
| "TransformersDepthLoader": TransformersDepthLoader, | |
| "Open3DReconstructionLoader": Open3DReconstructionLoader, | |
| "GaussianSplatLoader": GaussianSplatLoader, | |
| "DepthSplatLoader": DepthSplatLoader, | |
| } | |
| def build_loader(loader_name: str, model_id: str, device: torch.device, **kwargs) -> BaseLoader: | |
| """Instantiate a loader by name string.""" | |
| cls = _LOADER_MAP.get(loader_name) | |
| if cls is None: | |
| raise ValueError( | |
| f"Unknown loader: {loader_name!r}. " | |
| f"Available: {list(_LOADER_MAP)}" | |
| ) | |
| return cls(model_id=model_id, device=device, **kwargs) | |
| __all__ = [ | |
| "BaseLoader", | |
| "BiRefNetLoader", | |
| "TransformersDepthLoader", | |
| "Open3DReconstructionLoader", | |
| "GaussianSplatLoader", | |
| "DepthSplatLoader", | |
| "build_loader", | |
| ] | |