""" PixelDiT model loader. Usage: from modeling_pixeldit import load_pixeldit model = load_pixeldit() out = model(x, t, y) # [B,3,H,W], [B], [B,300,2304] -> [B,3,H,W] """ import sys import torch sys.path.insert(0, "/home/nobus/Raid0/PixelDiT") from pixdit_core.pixeldit_t2i import PixDiT_T2I _CKPT = ( "/home/nobus/.cache/huggingface/hub/" "models--nvidia--PixelDiT-1300M-1024px/snapshots/" "7c63b99a7a399918a1d6478b095698a65f664847/pixeldit_t2i_v1.pth" ) _ARCH = dict( in_channels=3, num_groups=24, hidden_size=1536, pixel_hidden_size=16, pixel_attn_hidden_size=1152, pixel_num_groups=16, patch_depth=14, pixel_depth=2, patch_size=16, txt_embed_dim=2304, txt_max_length=300, ) def load_pixeldit(checkpoint=_CKPT, device="cuda", dtype=torch.bfloat16): model = PixDiT_T2I(**_ARCH) state = torch.load(checkpoint, map_location="cpu", weights_only=False) sd = state.get("state_dict", state) sd = {(k[5:] if k.startswith("core.") else k): v for k, v in sd.items()} missing, _ = model.load_state_dict(sd, strict=False) if missing: print(f"[modeling] {len(missing)} missing keys (expected)") model = model.to(device).to(dtype).eval() print(f"[modeling] PixelDiT loaded — {sum(p.numel() for p in model.parameters()):,} params") return model