pixeldit-diffusers / pixeldit /modeling_pixeldit_hf.py
madtune's picture
Upload folder using huggingface_hub
e9200bf verified
Raw
History Blame
2.88 kB
"""
HF-compatible PixelDiT wrapper.
Allows save_pretrained / from_pretrained and peft LoRA targeting.
Usage:
# Convert from original .pth
model = PixelDiTModel.from_pth("pixeldit_t2i_v1.pth")
model.save_pretrained("pixeldit-diffusers/")
# Load back
model = PixelDiTModel.from_pretrained("pixeldit-diffusers/")
# LoRA
from peft import get_peft_model, LoraConfig
lora_cfg = LoraConfig(target_modules=["qkv_x", "qkv_y", "proj_x", "proj_y"])
model = get_peft_model(model, lora_cfg)
"""
import sys
import torch
from transformers import PreTrainedModel
sys.path.insert(0, "/home/nobus/Raid0/PixelDiT")
from pixdit_core.pixeldit_t2i import PixDiT_T2I
from .configuration_pixeldit import PixelDiTConfig
class PixelDiTModel(PreTrainedModel):
config_class = PixelDiTConfig
_tied_weights_keys = []
@property
def all_tied_weights_keys(self):
return {}
def __init__(self, config: PixelDiTConfig):
super().__init__(config)
self.model = PixDiT_T2I(
in_channels = config.in_channels,
num_groups = config.num_groups,
hidden_size = config.hidden_size,
pixel_hidden_size = config.pixel_hidden_size,
pixel_attn_hidden_size = config.pixel_attn_hidden_size,
pixel_num_groups = config.pixel_num_groups,
patch_depth = config.patch_depth,
pixel_depth = config.pixel_depth,
num_text_blocks = config.num_text_blocks,
patch_size = config.patch_size,
txt_embed_dim = config.txt_embed_dim,
txt_max_length = config.txt_max_length,
use_text_rope = config.use_text_rope,
text_rope_theta = config.text_rope_theta,
repa_encoder_index = config.repa_encoder_index,
use_pixel_abs_pos = config.use_pixel_abs_pos,
)
def forward(self, x, t, y, s=None, mask=None):
return self.model(x, t, y, s=s, mask=mask)
@classmethod
def from_pth(cls, pth_path: str, config: PixelDiTConfig = None):
"""Load from original nvidia .pth checkpoint, handles core. prefix."""
if config is None:
config = PixelDiTConfig()
model = cls(config)
state = torch.load(pth_path, map_location="cpu", weights_only=False)
sd = state.get("state_dict", state)
# strip trainer wrapper prefix, then add HF model. prefix
sd = {(k[5:] if k.startswith("core.") else k): v for k, v in sd.items()}
sd = {"model." + k: v for k, v in sd.items()}
missing, unexpected = model.load_state_dict(sd, strict=False)
print(f"[PixelDiTModel.from_pth] loaded — {len(missing)} missing, {len(unexpected)} unexpected")
return model