#!/usr/bin/env python3 """ export_text_enc.py — Export Qwen3-Embedding text encoder to ONNX. The text encoder is a standard Qwen3Model (28 layers, H=1024, causal attention) that takes BPE token IDs and produces hidden states for the condition encoder. Usage: python export_text_enc.py --model-dir --output Exports: text_encoder.onnx — Full 28-layer transformer Input: input_ids [B, S] int64 Output: hidden_states [B, S, 1024] fp16 embed_lookup.bin — Raw embedding table (vocab_size * hidden_size * 2 bytes, BF16) Used for lyric token embedding lookup on CPU (no ONNX needed). """ import argparse import os import sys import time import struct from pathlib import Path import numpy as np import torch import torch.nn as nn class TextEncoderWrapper(nn.Module): """Wrapper around Qwen3Model that returns hidden_states as a flat tensor. ONNX inputs: input_ids: [B, S] int64 — BPE token IDs ONNX output: hidden_states: [B, S, 1024] fp16 — last hidden state """ def __init__(self, model): super().__init__() self.model = model def forward(self, input_ids): outputs = self.model( input_ids=input_ids, attention_mask=None, # causal mask generated internally output_hidden_states=False, return_dict=True, ) return outputs.last_hidden_state def load_model(model_dir: str, device: str = "cuda", dtype=torch.float32): """Load Qwen3-Embedding model from safetensors.""" model_dir = Path(model_dir) # Fix Windows encoding issues if sys.platform == "win32": import io sys.stdout = io.TextIOWrapper(sys.stdout.buffer, encoding='utf-8', errors='replace') sys.stderr = io.TextIOWrapper(sys.stderr.buffer, encoding='utf-8', errors='replace') print(f"[export_text_enc] Loading model from {model_dir}...") t0 = time.time() from transformers import AutoModel, AutoConfig config = AutoConfig.from_pretrained(str(model_dir)) # Force SDPA for ONNX export (no flash attention) config._attn_implementation = "sdpa" model = AutoModel.from_pretrained( str(model_dir), config=config, torch_dtype=dtype, trust_remote_code=True, ) model = model.to(device) model.eval() t1 = time.time() n_params = sum(p.numel() for p in model.parameters()) / 1e6 print(f"[export_text_enc] Model loaded in {t1-t0:.1f}s ({n_params:.0f}M params)") print(f"[export_text_enc] Config: {config.num_hidden_layers}L, H={config.hidden_size}, " f"heads={config.num_attention_heads}/{config.num_key_value_heads}") return model, config def export_onnx(model, config, output_path: str, opset: int = 18): """Export the text encoder to ONNX.""" device = next(model.parameters()).device dtype = next(model.parameters()).dtype wrapper = TextEncoderWrapper(model) wrapper.eval() # Dummy inputs for tracing B = 1 S = 128 # typical sequence length dummy_input_ids = torch.randint(0, config.vocab_size, (B, S), device=device, dtype=torch.long) print(f"[export_text_enc] Tracing with shapes: input_ids={list(dummy_input_ids.shape)}") # Test forward pass print("[export_text_enc] Testing forward pass...") with torch.no_grad(): test_out = wrapper(dummy_input_ids) print(f"[export_text_enc] Output shape: {list(test_out.shape)} " f"(expected [{B}, {S}, {config.hidden_size}])") # Export to ONNX print(f"[export_text_enc] Exporting to ONNX (opset {opset})...") t0 = time.time() torch.onnx.export( wrapper, (dummy_input_ids,), output_path, opset_version=opset, input_names=["input_ids"], output_names=["hidden_states"], dynamic_axes={ "input_ids": {0: "batch", 1: "seq_len"}, "hidden_states": {0: "batch", 1: "seq_len"}, }, do_constant_folding=True, export_params=True, ) t1 = time.time() file_size = os.path.getsize(output_path) print(f"[export_text_enc] Exported to {output_path}") print(f"[export_text_enc] File size: {file_size/1e6:.1f} MB") print(f"[export_text_enc] Export time: {t1-t0:.1f}s") return output_path def export_embed_table(model, config, output_path: str): """Export the embedding table as a raw binary file for lyric lookup. The lyric path uses embed_tokens lookup only (no transformer layers). We export the table as float32 for direct CPU indexing. Format: raw float32 array [vocab_size, hidden_size] """ embed_weight = model.embed_tokens.weight.detach().cpu().float().numpy() V, H = embed_weight.shape with open(output_path, "wb") as f: # Header: vocab_size (int32), hidden_size (int32) f.write(struct.pack(" {output_path} ({file_size/1e6:.1f} MB)") def export_null_cond(model_dir: str, output_path: str): """Export null_condition_emb from the DiT model as raw float32. This is a [2048] float32 vector used for classifier-free guidance padding. Read from the DiT safetensors since it lives there. """ from safetensors.torch import load_file model_dir = Path(model_dir) st_path = model_dir / "model.safetensors" if not st_path.exists(): # Try multi-shard for p in sorted(model_dir.glob("model-*.safetensors")): st = load_file(str(p)) if "null_condition_emb" in st: vec = st["null_condition_emb"].detach().cpu().float().numpy() with open(output_path, "wb") as f: f.write(struct.pack(" {output_path}") return print("[export_text_enc] WARNING: null_condition_emb not found") return st = load_file(str(st_path)) if "null_condition_emb" not in st: print("[export_text_enc] WARNING: null_condition_emb not found in model.safetensors") return vec = st["null_condition_emb"].detach().cpu().float().numpy() with open(output_path, "wb") as f: f.write(struct.pack(" {output_path}") def verify_onnx(onnx_path: str, model, config): """Verify ONNX output matches PyTorch.""" try: import onnxruntime as ort except ImportError: print("[export_text_enc] onnxruntime not installed, skipping verification") return device = next(model.parameters()).device wrapper = TextEncoderWrapper(model) wrapper.eval() # Test inputs B, S = 1, 64 input_ids = torch.randint(0, config.vocab_size, (B, S), device=device, dtype=torch.long) # PyTorch reference with torch.no_grad(): ref_out = wrapper(input_ids).cpu().float().numpy() # ONNX inference providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] sess = ort.InferenceSession(onnx_path, providers=providers) ort_out = sess.run(None, { "input_ids": input_ids.cpu().numpy(), })[0] # Compare max_diff = np.max(np.abs(ref_out - ort_out)) mean_diff = np.mean(np.abs(ref_out - ort_out)) print(f"[export_text_enc] Verification: max_diff={max_diff:.6f}, mean_diff={mean_diff:.6f}") if max_diff < 0.05: print("[export_text_enc] PASS: ONNX output matches PyTorch (within FP16 tolerance)") else: print("[export_text_enc] WARNING: Large difference — may need investigation") def main(): parser = argparse.ArgumentParser(description="Export Qwen3-Embedding text encoder to ONNX") parser.add_argument("--model-dir", required=True, help="Path to Qwen3-Embedding-0.6B directory") parser.add_argument("--output", default=None, help="Output ONNX file (default: models/onnx/text_encoder.onnx)") parser.add_argument("--dit-dir", default=None, help="Path to DiT model dir (for null_condition_emb export)") parser.add_argument("--opset", type=int, default=18, help="ONNX opset version (default: 18)") parser.add_argument("--verify", action="store_true", help="Verify ONNX output matches PyTorch") parser.add_argument("--device", default="cuda", help="Device for model loading (default: cuda)") args = parser.parse_args() # Default output path if args.output is None: onnx_dir = Path(args.model_dir).parent / "onnx" onnx_dir.mkdir(parents=True, exist_ok=True) args.output = str(onnx_dir / "text_encoder.onnx") os.makedirs(os.path.dirname(args.output), exist_ok=True) output_dir = os.path.dirname(args.output) # Load model model, config = load_model(args.model_dir, device=args.device) # Export ONNX export_onnx(model, config, args.output, opset=args.opset) # Export embedding table for lyric lookup embed_path = os.path.join(output_dir, "embed_tokens.bin") export_embed_table(model, config, embed_path) # Export null_condition_emb if DiT dir provided if args.dit_dir: null_cond_path = os.path.join(output_dir, "null_condition_emb.bin") export_null_cond(args.dit_dir, null_cond_path) # Verify if args.verify: verify_onnx(args.output, model, config) print("[export_text_enc] Done!") if __name__ == "__main__": main()