Files
2026-08-16 18:24:52 +07:00

288 lines
9.8 KiB
Python

#!/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 <path-to-Qwen3-Embedding-0.6B> --output <output.onnx>
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("<II", V, H))
# Raw float32 weights
f.write(embed_weight.tobytes())
file_size = os.path.getsize(output_path)
print(f"[export_text_enc] Embedding table: [{V}, {H}] -> {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("<I", vec.shape[0]))
f.write(vec.tobytes())
print(f"[export_text_enc] null_condition_emb: [{vec.shape[0]}] -> {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("<I", vec.shape[0]))
f.write(vec.tobytes())
print(f"[export_text_enc] null_condition_emb: [{vec.shape[0]}] -> {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()