omni-diffusion / magvit /modeling_utils.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
fe78f25 verified
Raw
History Blame Contribute Delete
1 kB
"""Minimal shims for the MagViT-v2 modules.
The upstream MagViT code (from MMaDA) subclasses diffusers' `ModelMixin` /
`ConfigMixin` and uses `@register_to_config`. For inference we only need
`nn.Module` behaviour plus a no-op config recorder, so we provide lightweight
replacements and avoid depending on diffusers' fast-moving internal loading
APIs. Weights are loaded directly from safetensors in `load_magvit.py`.
"""
import functools
import torch.nn as nn
class ConfigMixin:
config_name = "config.json"
class ModelMixin(nn.Module):
"""Standin for diffusers.ModelMixin — just an nn.Module for inference."""
def __init__(self, *args, **kwargs):
super().__init__()
def register_to_config(init):
"""No-op replacement for diffusers' register_to_config decorator."""
@functools.wraps(init)
def inner(self, *args, **kwargs):
if not hasattr(self, "config"):
self.config = dict(kwargs)
init(self, *args, **kwargs)
return inner