Download load_model.py from dejanseo/fanout-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 9.7 kB
-
https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/load_model.py
- Command line
-
hf download hf://dejanseo/fanout-diffusion/load_model.py
-
curl -L -o load_model.py https://huggingface.co/dejanseo/fanout-diffusion/resolve/main/load_model.py
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!") | |