fanout-diffusion / load_model.py
dejanseo's picture
Fix ONNX input/output feed tensor names
6a539fa verified
Raw History Blame Contribute Delete
9.7 kB
"""Load fanout-diffusion models and datasets from Hugging Face Hub.
Repository: dejanseo/fanout-diffusion (Private)
Requires: huggingface_hub, torch, onnxruntime (optional for ONNX)
"""
from __future__ import annotations
import argparse
import os
from pathlib import Path
from typing import Any, Dict, Optional, Tuple, Union
import torch
from huggingface_hub import hf_hub_download
from torch import Tensor
# ---------------------------------------------------------------------------
# 1. New 1-Step Consistency Model (ONNX Runtime - Recommended / Portable)
# ---------------------------------------------------------------------------
def load_fanout_onnx(
repo_id: str = "dejanseo/fanout-diffusion",
filename: str = "checkpoints/champion_b1_consistency_1step.onnx",
use_cuda: bool = True,
token: Optional[str] = None,
):
"""
Download and initialize the ONNX Runtime InferenceSession for 1-step consistency fanout.
Requires: pip install onnxruntime (or onnxruntime-gpu).
"""
import onnxruntime as ort
print(f"Fetching ONNX model '{filename}' from '{repo_id}'...")
model_path = hf_hub_download(
repo_id=repo_id,
filename=filename,
repo_type="model",
token=token or os.environ.get("HF_TOKEN"),
)
providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] if use_cuda else ["CPUExecutionProvider"]
session = ort.InferenceSession(model_path, providers=providers)
print(f"Loaded ONNX InferenceSession (Provider: {session.get_providers()[0]})")
return session
def sample_consistency_onnx(session, query_vector: torch.Tensor | Any) -> torch.Tensor:
"""
Generate 10 continuous fanout vectors from a query embedding in a single forward pass.
query_vector: [B, 768] (float32, normalized)
Returns: [B, 10, 768]
"""
import numpy as np
if isinstance(query_vector, torch.Tensor):
q_np = query_vector.detach().cpu().numpy().astype(np.float32)
else:
q_np = np.asarray(query_vector, dtype=np.float32)
if q_np.ndim == 1:
q_np = q_np[None, :]
B = q_np.shape[0]
# Noise: sigma_max = 80.0
noise_np = (np.random.randn(B, 10, 768) * 80.0).astype(np.float32)
sigmas_np = np.full((B,), 80.0, dtype=np.float32)
outputs = session.run(
["denoised_fanout"],
{"noisy_slots": noise_np, "sigma": sigmas_np, "query_emb": q_np},
)
return torch.from_numpy(outputs[0])
# ---------------------------------------------------------------------------
# 2. PyTorch 1-Step Consistency Distillation Model
# ---------------------------------------------------------------------------
def load_fanout_consistency(
repo_id: str = "dejanseo/fanout-diffusion",
checkpoint_filename: str = "checkpoints/champion_b1_consistency_1step_qat.pt",
device: Optional[Union[str, torch.device]] = None,
token: Optional[str] = None,
):
"""
Download and initialize the PyTorch 1-step consistency denoiser (B1EDMDenoiser).
"""
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(device)
print(f"Fetching consistency checkpoint '{checkpoint_filename}' from '{repo_id}'...")
ckpt_path = hf_hub_download(
repo_id=repo_id,
filename=checkpoint_filename,
repo_type="model",
token=token or os.environ.get("HF_TOKEN"),
)
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
config = ckpt["config"]
try:
from src.r4t.b1_diffusion import B1EDMDenoiser
from scripts.quantize_outer_int4 import unpack_int4_signed
except ImportError:
raise ImportError("src.r4t modules required to load PyTorch B1EDMDenoiser checkpoint.")
model = B1EDMDenoiser(config, backend="tc" if torch.cuda.is_available() else "cpu", pure_1bit=False).to(device)
model.freeze_for_inference()
state = model.state_dict()
if "weights" in ckpt:
for k, v in ckpt["weights"].items():
if k in state:
state[k].copy_(v.to(device))
if "int4_outer" in ckpt:
for k, d in ckpt["int4_outer"].items():
if k in state:
state[k].copy_(unpack_int4_signed(d["packed"].to(device), d["scale"].to(device)))
elif "model_state_dict" in ckpt:
model.load_state_dict(ckpt["model_state_dict"], strict=False)
model.eval()
print(f"Loaded 1-Step Consistency Model (1-Bit TC + INT4 Outer, Device: {device})")
return model, config
def sample_consistency_pytorch(model, query_vector: torch.Tensor) -> torch.Tensor:
"""Generate 10 continuous fanout vectors in 1 single forward pass."""
device = next(model.parameters()).device
q = query_vector.to(device).float()
if q.ndim == 1:
q = q.unsqueeze(0)
B = q.size(0)
sigma_max = getattr(model.config, "sigma_max", 80.0)
noise = torch.randn(B, model.config.sequence_length, model.config.embedding_dim, device=device) * sigma_max
sigmas = torch.full((B,), sigma_max, device=device)
with torch.no_grad():
out = model(noise, sigmas, q)
return out
# ---------------------------------------------------------------------------
# 3. Legacy Multi-Step EDM Model (Backward Compatibility)
# ---------------------------------------------------------------------------
def load_fanout_diffusion(
repo_id: str = "dejanseo/fanout-diffusion",
checkpoint_filename: str = "diffusion_540k_best.pt",
device: Optional[Union[str, torch.device]] = None,
token: Optional[str] = None,
):
"""Download and instantiate legacy multi-step EDMDenoiser from Hugging Face Hub."""
if device is None:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
else:
device = torch.device(device)
print(f"Fetching '{checkpoint_filename}' from '{repo_id}'...")
checkpoint_path = hf_hub_download(
repo_id=repo_id,
filename=checkpoint_filename,
repo_type="model",
token=token or os.environ.get("HF_TOKEN"),
)
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
config = checkpoint["config"]
from src.r4t.diffusion import EDMDenoiser
model = EDMDenoiser(config).to(device)
if "ema_state_dict" in checkpoint and "shadow" in checkpoint["ema_state_dict"]:
model.load_state_dict(checkpoint["ema_state_dict"]["shadow"])
elif "model_state_dict" in checkpoint:
model.load_state_dict(checkpoint["model_state_dict"])
else:
raise KeyError("No valid state_dict found in checkpoint.")
model.eval()
val_loss = checkpoint.get("val_loss", None)
epoch = checkpoint.get("epoch", None)
loss_str = f"{val_loss:.4f}" if val_loss is not None else "N/A"
print(f"Loaded Legacy EDMDenoiser (Epoch: {epoch}, Val Loss: {loss_str}, Device: {device})")
return model, config
# ---------------------------------------------------------------------------
# 4. Dataset Loader
# ---------------------------------------------------------------------------
def load_fanout_dataset(
repo_id: str = "dejanseo/fanout-diffusion",
filename: str = "data/diffusion_dataset_540k.pt",
token: Optional[str] = None,
) -> dict:
"""Download and load the pre-computed embedding dataset from Hugging Face Hub."""
print(f"Fetching dataset '{filename}' from '{repo_id}'...")
dataset_path = hf_hub_download(
repo_id=repo_id,
filename=filename,
repo_type="model",
token=token or os.environ.get("HF_TOKEN"),
)
data = torch.load(dataset_path, map_location="cpu", weights_only=False)
n = data.get("query_embeddings", data.get("queries")).shape[0]
print(f"Loaded dataset: {n:,} query sets")
return data
# ---------------------------------------------------------------------------
# CLI Test Runner
# ---------------------------------------------------------------------------
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="Load and test fanout-diffusion models from Hugging Face Hub")
parser.add_argument("--repo_id", type=str, default="dejanseo/fanout-diffusion")
parser.add_argument("--mode", type=str, choices=["onnx", "consistency", "legacy"], default="onnx")
parser.add_argument("--query", type=str, default="running shoes and athletic sneakers")
args = parser.parse_args()
print(f"--- Testing {args.mode.upper()} mode with query: '{args.query}' ---")
if args.mode == "onnx":
session = load_fanout_onnx(repo_id=args.repo_id)
dummy_q = torch.randn(1, 768)
dummy_q = torch.nn.functional.normalize(dummy_q, p=2, dim=-1)
res = sample_consistency_onnx(session, dummy_q)
print("Generated shape:", res.shape)
print("Inference verified successfully!")
elif args.mode == "consistency":
model, config = load_fanout_consistency(repo_id=args.repo_id)
dummy_q = torch.randn(1, 768)
dummy_q = torch.nn.functional.normalize(dummy_q, p=2, dim=-1)
res = sample_consistency_pytorch(model, dummy_q)
print("Generated shape:", res.shape)
print("Inference verified successfully!")
elif args.mode == "legacy":
from src.r4t.diffusion import sample_edm
model, config = load_fanout_diffusion(repo_id=args.repo_id)
dummy_q = torch.randn(1, 768, device=next(model.parameters()).device)
dummy_q = torch.nn.functional.normalize(dummy_q, p=2, dim=-1)
res = sample_edm(model, dummy_q, sampling_steps=16)
print("Generated shape:", res.shape)
print("Inference verified successfully!")