"""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!")