| |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| """ |
| Export GR00T N1.7 model components to ONNX for TensorRT optimization. |
| |
| Supports three export modes: |
| - dit_only: Export only the DiT (backward compatible with N1.6). |
| - action_head: Export 4 action head components (ViT + LLM stay in PyTorch). |
| - full_pipeline: Export ViT + LLM + 4 action head components. Lightweight |
| glue ops (embed_tokens, masked_scatter, get_rope_index, |
| VLLN) remain in PyTorch. Referred to as 'n17_full_pipeline' |
| by the engine-loading code (trt_model_forward.py, |
| build_trt_pipeline.py) and as 'trt_full_pipeline' by the |
| standalone_inference_script.py --inference-mode flag; |
| both names describe the same engine set. |
| |
| Usage: |
| # Download finetuned model first (HF doesn't support nested repo paths) |
| uv run hf download nvidia/GR00T-N1.7-LIBERO --include "libero_10/config.json" "libero_10/embodiment_id.json" "libero_10/model-*.safetensors" "libero_10/model.safetensors.index.json" "libero_10/processor_config.json" "libero_10/statistics.json" --local-dir checkpoints/GR00T-N1.7-LIBERO |
| |
| # DiT only (default). N1.7 TRT export uses the legacy ONNX exporter |
| # explicitly (`dynamo=False`) so dynamic axes remain TensorRT-friendly. |
| python export_onnx_n1d7.py \\ |
| --model-path checkpoints/GR00T-N1.7-LIBERO/libero_10 \\ |
| --dataset-path demo_data/libero_demo \\ |
| --output-dir ./gr00t_trt_deployment/onnx |
| |
| # Action head (4 components) |
| python export_onnx_n1d7.py \\ |
| --model-path checkpoints/GR00T-N1.7-LIBERO/libero_10 \\ |
| --dataset-path demo_data/libero_demo \\ |
| --output-dir ./gr00t_trt_deployment/onnx \\ |
| --export-mode action_head |
| """ |
|
|
| import copy |
| from dataclasses import dataclass |
| import json |
| import logging |
| import os |
| from pathlib import Path |
| from typing import Literal, Optional |
|
|
| from _trt_contract import EXPORT_METADATA_SCHEMA_VERSION |
| from gr00t.data.dataset.lerobot_episode_loader import LeRobotEpisodeLoader |
| from gr00t.data.dataset.sharded_single_step_dataset import extract_step_data |
| from gr00t.data.embodiment_tags import EmbodimentTag |
| from gr00t.data.utils import parse_observation_gr00t |
| from gr00t.deployment.modes import ExportMode |
| from gr00t.model.modules.qwen3_backbone import _assign_inv_freq, recompute_vision_rotary_inv_freq |
| from gr00t.policy.gr00t_policy import Gr00tPolicy |
| import numpy as np |
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| import torch.onnx |
| import tyro |
|
|
|
|
| |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s") |
| logger = logging.getLogger(__name__) |
|
|
|
|
| def _consolidate_external_data(onnx_path: str) -> None: |
| """Merge scattered external-data files into a single .data file next to the ONNX model.""" |
| import onnx |
| from onnx.external_data_helper import convert_model_to_external_data |
|
|
| onnx_dir = os.path.dirname(onnx_path) |
| onnx_name = os.path.basename(onnx_path) |
| data_file = onnx_name + ".data" |
|
|
| |
| scattered = [ |
| f |
| for f in os.listdir(onnx_dir) |
| if os.path.isfile(os.path.join(onnx_dir, f)) |
| and not f.endswith((".onnx", ".json", ".data")) |
| and f != data_file |
| ] |
| if not scattered: |
| return |
|
|
| logger.info(f" Consolidating {len(scattered)} external data files into {data_file}...") |
| model = onnx.load(onnx_path, load_external_data=True) |
| convert_model_to_external_data( |
| model, all_tensors_to_one_file=True, location=data_file, size_threshold=0 |
| ) |
| onnx.save(model, onnx_path) |
|
|
| |
| for f in scattered: |
| os.remove(os.path.join(onnx_dir, f)) |
| logger.info(f" Consolidated and cleaned up {len(scattered)} files.") |
|
|
|
|
| def verify_onnx_export(onnx_path: str) -> None: |
| """Load and check the exported ONNX model for validity.""" |
| import onnx |
|
|
| logger.info(f" Verifying {onnx_path} ...") |
| onnx.checker.check_model(onnx_path) |
| logger.info(" ONNX model verified successfully.") |
|
|
|
|
| def verify_onnx_with_ort( |
| onnx_path: str, |
| pytorch_module: torch.nn.Module, |
| sample_inputs: dict[str, torch.Tensor], |
| output_names: list[str], |
| label: str = "model", |
| ) -> dict[str, float]: |
| """Run ONNX Runtime inference and compare against PyTorch output. |
| |
| Returns dict of {output_name: cosine_similarity}. |
| Requires onnxruntime-gpu; skips gracefully if not installed. |
| """ |
| try: |
| import onnxruntime as ort |
| except ImportError: |
| logger.warning(" onnxruntime not installed — skipping ORT verification") |
| return {} |
|
|
| logger.info(f" ORT verification for {label}...") |
|
|
| |
| with torch.inference_mode(): |
| pt_inputs = tuple(sample_inputs[name] for name in sample_inputs) |
| pt_outputs = pytorch_module(*pt_inputs) |
| if not isinstance(pt_outputs, (tuple, list)): |
| pt_outputs = (pt_outputs,) |
|
|
| |
| providers = ["CUDAExecutionProvider", "CPUExecutionProvider"] |
| sess = ort.InferenceSession(onnx_path, providers=providers) |
| ort_inputs = {name: t.cpu().numpy() for name, t in sample_inputs.items()} |
| ort_outputs = sess.run(output_names, ort_inputs) |
|
|
| |
| results = {} |
| for i, name in enumerate(output_names): |
| pt_flat = pt_outputs[i].float().flatten().cpu() |
| ort_flat = torch.tensor(ort_outputs[i]).float().flatten() |
| cosine = torch.nn.functional.cosine_similarity( |
| pt_flat.unsqueeze(0), ort_flat.unsqueeze(0) |
| ).item() |
| results[name] = cosine |
| logger.info(f" {name}: ORT vs PyTorch cosine = {cosine:.6f}") |
|
|
| return results |
|
|
|
|
| |
| |
| |
|
|
|
|
| class DiTInputCapture: |
| """Capture DiT forward pass inputs during inference via a pre-forward hook.""" |
|
|
| def __init__(self): |
| self.captured = False |
| self.sa_embs = None |
| self.vl_embs = None |
| self.timestep = None |
| self.image_mask = None |
| self.backbone_attention_mask = None |
|
|
| def hook_fn(self, module, args, kwargs): |
| """Pre-forward hook to capture inputs.""" |
| if not self.captured: |
| self.sa_embs = kwargs["hidden_states"].detach().cpu().clone() |
| self.vl_embs = kwargs["encoder_hidden_states"].detach().cpu().clone() |
| self.timestep = kwargs["timestep"].detach().cpu().clone() |
| i_mask = kwargs.get("image_mask") |
| if i_mask is not None: |
| self.image_mask = i_mask.detach().cpu().clone() |
| bb_mask = kwargs.get("backbone_attention_mask") |
| if bb_mask is not None: |
| self.backbone_attention_mask = bb_mask.detach().cpu().clone() |
|
|
| self.captured = True |
| logger.info(" Captured DiT inputs:") |
| logger.info(f" sa_embs: {self.sa_embs.shape}") |
| logger.info(f" vl_embs: {self.vl_embs.shape}") |
| logger.info(f" timestep: {self.timestep.shape}") |
| if self.image_mask is not None: |
| logger.info(f" image_mask: {self.image_mask.shape}") |
| if self.backbone_attention_mask is not None: |
| logger.info(f" backbone_attention_mask: {self.backbone_attention_mask.shape}") |
|
|
|
|
| class ViTInputCapture: |
| """Capture ViT (VisionModel) forward inputs/outputs during inference.""" |
|
|
| def __init__(self): |
| self.captured = False |
| self.pixel_values_shape = None |
| self.grid_thw = None |
| self.output_shape = None |
| self.deepstack_shapes = [] |
|
|
| def hook_fn(self, module, args, kwargs, output): |
| if not self.captured: |
| self.pixel_values_shape = args[0].shape |
| grid = args[1] if len(args) > 1 else kwargs.get("grid_thw") |
| self.grid_thw = grid.detach().cpu().clone() |
| if isinstance(output, tuple): |
| self.output_shape = output[0].shape |
| self.deepstack_shapes = [f.shape for f in output[1]] |
| else: |
| self.output_shape = output.shape |
| self.captured = True |
| logger.info(" Captured ViT inputs:") |
| logger.info(f" pixel_values: {self.pixel_values_shape}") |
| logger.info(f" grid_thw: {self.grid_thw.tolist()}") |
| logger.info(f" output: {self.output_shape}") |
| logger.info(f" deepstack: {len(self.deepstack_shapes)} features") |
|
|
|
|
| class LLMInputCapture: |
| """Capture LLM (Qwen3VLTextModel) inputs during inference via a pre-forward hook. |
| |
| Captures inputs_embeds, position_ids, attention_mask, visual_pos_masks, |
| and deepstack_visual_embeds — everything needed to reproduce the LLM forward. |
| """ |
|
|
| def __init__(self): |
| self.captured = False |
| self.inputs_embeds = None |
| self.position_ids = None |
| self.attention_mask = None |
| self.visual_pos_masks = None |
| self.deepstack_visual_embeds = None |
|
|
| def hook_fn(self, module, args, kwargs): |
| if not self.captured: |
| ie = kwargs.get("inputs_embeds") |
| if ie is None and len(args) > 0: |
| ie = args[0] |
| if ie is not None: |
| self.inputs_embeds = ie.detach().cpu().clone() |
|
|
| pid = kwargs.get("position_ids") |
| if pid is not None: |
| self.position_ids = pid.detach().cpu().clone() |
|
|
| am = kwargs.get("attention_mask") |
| if am is not None: |
| self.attention_mask = am.detach().cpu().clone() |
|
|
| vpm = kwargs.get("visual_pos_masks") |
| if vpm is not None: |
| self.visual_pos_masks = vpm.detach().cpu().clone() |
|
|
| dve = kwargs.get("deepstack_visual_embeds") |
| if dve is not None: |
| self.deepstack_visual_embeds = [d.detach().cpu().clone() for d in dve] |
|
|
| self.captured = True |
| logger.info(" Captured LLM inputs:") |
| if self.inputs_embeds is not None: |
| logger.info(f" inputs_embeds: {self.inputs_embeds.shape}") |
| if self.position_ids is not None: |
| logger.info(f" position_ids: {self.position_ids.shape}") |
| if self.attention_mask is not None: |
| logger.info(f" attention_mask: {self.attention_mask.shape}") |
| if self.visual_pos_masks is not None: |
| logger.info(f" visual_pos_masks: {self.visual_pos_masks.shape}") |
| if self.deepstack_visual_embeds is not None: |
| logger.info( |
| f" deepstack: {len(self.deepstack_visual_embeds)} tensors, " |
| f"shapes: {[d.shape for d in self.deepstack_visual_embeds]}" |
| ) |
|
|
|
|
| |
| |
| |
|
|
|
|
| def _apply_rotary_real(x, cos, sin): |
| """Apply rotary position embeddings using only real-valued ops (no complex). |
| |
| Uses float32 internally to match transformers' apply_rotary_pos_emb_vision |
| precision, then casts back to the original dtype. |
| |
| Args: |
| x: [seq, heads, head_dim] |
| cos, sin: [seq, head_dim] |
| Returns: |
| [seq, heads, head_dim] |
| """ |
| orig_dtype = x.dtype |
| x = x.float() |
| cos = cos.float().unsqueeze(1) |
| sin = sin.float().unsqueeze(1) |
| half = x.shape[-1] // 2 |
| x1 = x[..., :half] |
| x2 = x[..., half:] |
| rotated = torch.cat((-x2, x1), dim=-1) |
| return (x * cos + rotated * sin).to(orig_dtype) |
|
|
|
|
| def _make_onnx_vision_attention_forward(attn_module, chunk_sizes=None): |
| """Create an ONNX-exportable attention forward for a single VisionAttention. |
| |
| Three key changes from the original: |
| 1. Replaces cu_seqlens-based splitting with static chunk splitting |
| (each image's patches attend only within their own chunk) |
| 2. Replaces apply_rotary_pos_emb_vision (uses complex numbers) with |
| real-valued rotate_half implementation |
| 3. Casts to float32 before softmax for TRT accuracy |
| |
| Args: |
| attn_module: The VisionAttention module to wrap |
| chunk_sizes: List of ints, number of patches per image. |
| e.g. [256, 256] for 2 images of 256 patches each. |
| If None or single chunk, does full-sequence attention. |
| """ |
|
|
| def forward( |
| hidden_states, cu_seqlens=None, rotary_pos_emb=None, position_embeddings=None, **kwargs |
| ): |
| seq_length = hidden_states.shape[0] |
| qkv = attn_module.qkv(hidden_states) |
| qkv = qkv.reshape(seq_length, 3, attn_module.num_heads, -1) |
| qkv = qkv.permute(1, 0, 2, 3) |
| q, k, v = qkv.unbind(0) |
| |
|
|
| cos, sin = position_embeddings |
| q = _apply_rotary_real(q, cos, sin) |
| k = _apply_rotary_real(k, cos, sin) |
|
|
| |
| |
| if chunk_sizes is not None and len(chunk_sizes) > 1: |
| q_chunks = torch.split(q, chunk_sizes, dim=0) |
| k_chunks = torch.split(k, chunk_sizes, dim=0) |
| v_chunks = torch.split(v, chunk_sizes, dim=0) |
|
|
| attn_outputs = [] |
| for q_c, k_c, v_c in zip(q_chunks, k_chunks, v_chunks): |
| |
| q_c = q_c.transpose(0, 1) |
| k_c = k_c.transpose(0, 1) |
| v_c = v_c.transpose(0, 1) |
|
|
| w = torch.matmul(q_c, k_c.transpose(-2, -1)) * attn_module.scaling |
| w = w.to(torch.float32) |
| w = F.softmax(w, dim=-1) |
| w = w.to(v_c.dtype) |
| out = torch.matmul(w, v_c) |
| attn_outputs.append(out.transpose(0, 1)) |
|
|
| attn_output = torch.cat(attn_outputs, dim=0) |
| else: |
| |
| q = q.transpose(0, 1) |
| k = k.transpose(0, 1) |
| v = v.transpose(0, 1) |
|
|
| attn_weights = torch.matmul(q, k.transpose(-2, -1)) * attn_module.scaling |
| attn_weights = attn_weights.to(torch.float32) |
| attn_weights = F.softmax(attn_weights, dim=-1) |
| attn_weights = attn_weights.to(v.dtype) |
| attn_output = torch.matmul(attn_weights, v) |
| attn_output = attn_output.transpose(0, 1) |
|
|
| |
| attn_output = attn_output.reshape(seq_length, -1).contiguous() |
| attn_output = attn_module.proj(attn_output) |
| return attn_output |
|
|
| return forward |
|
|
|
|
| def _patch_vision_attention_for_export(vision_model, chunk_sizes=None): |
| """Monkey-patch all VisionAttention.forward with ONNX-friendly versions. |
| |
| Args: |
| vision_model: The vision model whose attention blocks to patch |
| chunk_sizes: List of ints, patches per image (from grid_thw). |
| Enables per-image attention splitting. If None, full-sequence attention. |
| |
| Returns list of original forwards for restoration after export. |
| """ |
| originals = [] |
| for block in vision_model.blocks: |
| attn = block.attn |
| originals.append(attn.forward) |
| attn.forward = _make_onnx_vision_attention_forward(attn, chunk_sizes=chunk_sizes) |
| info = f" Patched {len(originals)} vision attention blocks for ONNX export" |
| if chunk_sizes and len(chunk_sizes) > 1: |
| info += f" (chunk_sizes={chunk_sizes})" |
| logger.info(info) |
| return originals |
|
|
|
|
| def _restore_vision_attention(vision_model, originals): |
| """Restore original attention forwards after export.""" |
| for block, orig in zip(vision_model.blocks, originals): |
| block.attn.forward = orig |
|
|
|
|
| class Qwen3VisionForExport(torch.nn.Module): |
| """ONNX-exportable wrapper for Qwen3-VL Vision Model. |
| |
| Pre-computes position embeddings and rotary embeddings for a fixed grid_thw |
| to avoid ComplexDouble operations that ONNX cannot handle. Replaces the |
| dynamic VisionModel.forward with a traceable version. |
| |
| Architecture: patch_embed → add pos_embed → blocks(attn+ffn) → deepstack → merger |
| """ |
|
|
| def __init__(self, vision_model, grid_thw: torch.Tensor): |
| super().__init__() |
| self.patch_embed = vision_model.patch_embed |
| self.blocks = vision_model.blocks |
| self.merger = vision_model.merger |
| self.deepstack_visual_indexes = vision_model.deepstack_visual_indexes |
| self.deepstack_merger_list = vision_model.deepstack_merger_list |
|
|
| |
| |
| with torch.no_grad(): |
| pos_embeds = vision_model.fast_pos_embed_interpolate(grid_thw) |
| rotary = vision_model.rot_pos_emb(grid_thw) |
| emb = torch.cat((rotary, rotary), dim=-1) |
|
|
| self.register_buffer("_pos_embeds", pos_embeds.clone().detach().contiguous()) |
| self.register_buffer("_rot_cos", emb.cos().clone().detach().contiguous()) |
| self.register_buffer("_rot_sin", emb.sin().clone().detach().contiguous()) |
|
|
| def forward(self, pixel_values): |
| hidden_states = self.patch_embed(pixel_values) |
| hidden_states = hidden_states + self._pos_embeds |
|
|
| position_embeddings = (self._rot_cos, self._rot_sin) |
|
|
| deepstack_features = [] |
| for layer_num, blk in enumerate(self.blocks): |
| hidden_states = blk( |
| hidden_states, |
| cu_seqlens=None, |
| position_embeddings=position_embeddings, |
| ) |
| if layer_num in self.deepstack_visual_indexes: |
| idx = self.deepstack_visual_indexes.index(layer_num) |
| deepstack_features.append(self.deepstack_merger_list[idx](hidden_states)) |
|
|
| image_embeds = self.merger(hidden_states) |
|
|
| if deepstack_features: |
| deepstack = torch.stack(deepstack_features) |
| else: |
| deepstack = image_embeds.new_zeros(1, 1, 1) |
|
|
| return image_embeds, deepstack |
|
|
|
|
| |
| |
| |
|
|
|
|
| def _restore_vision_rotary_inv_freq_fp32(vision, vision_config): |
| """Re-derive the ViT RoPE ``inv_freq`` in fp32 right before baking it into the engine. |
| |
| ``Gr00tPolicy`` loads the whole model in bf16 (``gr00t_policy.py``), which rounds |
| this non-persistent buffer even though the ViT itself is exported in fp32 for TRT |
| accuracy (max|delta| ~1.8e-4 for head_dim=64 -- a pure bf16 round-trip error). |
| Re-deriving the analytic fp32 value here keeps the baked rotary cos/sin at full |
| fp32 precision and lets the analytic-oracle guard verify a clean buffer instead of |
| the bf16-degraded one. Fails closed when the rotary submodule/layout is missing. |
| """ |
| rotary = getattr(vision, "rotary_pos_emb", None) |
| if rotary is None or not hasattr(rotary, "inv_freq"): |
| raise RuntimeError( |
| "ViT export: vision rotary_pos_emb/inv_freq not found; cannot rebuild the " |
| "rotary buffers before baking them into the engine (transformers Qwen3-VL " |
| "layout drift). Refusing to export unverified rotary." |
| ) |
| head_dim = vision_config.hidden_size // vision_config.num_heads |
| fp32_inv_freq = recompute_vision_rotary_inv_freq(rotary, head_dim // 2, rotary.inv_freq.device) |
| _assign_inv_freq(rotary, "inv_freq", fp32_inv_freq, persistent=False) |
|
|
|
|
| def _assert_vision_rotary_matches_analytic(vision, vision_config, *, atol=1e-5, rtol=1e-4): |
| """Abort export if the loaded ViT RoPE ``inv_freq`` drifts from the analytic oracle. |
| |
| The baked rotary cos/sin derive from this non-persistent buffer; the post-export |
| cosine checks are backend-vs-backend and cannot catch a common-mode error in it. |
| """ |
| rotary = getattr(vision, "rotary_pos_emb", None) |
| if rotary is None or not hasattr(rotary, "inv_freq"): |
| raise RuntimeError( |
| "ViT export: vision rotary_pos_emb/inv_freq not found; cannot verify the " |
| "rotary buffers before baking them into the engine (transformers Qwen3-VL " |
| "layout drift). Refusing to export unverified rotary." |
| ) |
| head_dim = vision_config.hidden_size // vision_config.num_heads |
| baked = rotary.inv_freq.detach().float() |
| analytic = recompute_vision_rotary_inv_freq(rotary, head_dim // 2, baked.device) |
| if not torch.allclose(baked, analytic, atol=atol, rtol=rtol): |
| max_err = (baked - analytic).abs().max().item() |
| raise RuntimeError( |
| "ViT export: loaded vision RoPE inv_freq diverges from the analytic oracle " |
| f"(max|delta|={max_err:.3e}); the baked rotary cos/sin would be silently wrong " |
| "and the backend-vs-backend cosine checks cannot catch it. Ensure " |
| "Qwen3Backbone._reset_rotary_inv_freq ran at load before exporting." |
| ) |
|
|
|
|
| def export_vit_to_onnx(policy, output_dir, captured_vit, use_bf16=True, batch_size=1): |
| """Export Qwen3-VL Vision Model to ONNX. |
| |
| Pre-computes position/rotary embeddings for the captured grid_thw to avoid |
| ComplexDouble ops. Monkey-patches attention to use standard SDPA (valid for |
| single-image inference where all patches attend to all patches). |
| |
| Input: pixel_values [num_patches * batch_size, C*T*pH*pW] |
| Output: image_embeds [num_merged_patches * batch_size, hidden_dim], |
| deepstack_features [num_layers, num_merged_patches * batch_size, hidden_dim] |
| """ |
| logger.info("\n" + "=" * 80) |
| logger.info("Exporting ViT (Qwen3-VL Vision) to ONNX") |
| logger.info("=" * 80) |
|
|
| backbone = policy.model.backbone |
| qwen_model = backbone.model |
| vision = qwen_model.model.visual |
|
|
| |
| |
| |
| |
| _restore_vision_rotary_inv_freq_fp32(vision, qwen_model.config.vision_config) |
| _assert_vision_rotary_matches_analytic(vision, qwen_model.config.vision_config) |
|
|
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
| vision = vision.to(dtype).eval().cuda() |
|
|
| |
| |
| |
| grid_thw = captured_vit.grid_thw.to(device="cuda") |
| |
| if batch_size > 1: |
| grid_thw = grid_thw.repeat(batch_size, 1) |
| chunk_sizes = torch.repeat_interleave(grid_thw[:, 1] * grid_thw[:, 2], grid_thw[:, 0]).tolist() |
| logger.info(f" grid_thw: {grid_thw.tolist()}, chunk_sizes: {chunk_sizes}") |
|
|
| |
| originals = _patch_vision_attention_for_export(vision, chunk_sizes=chunk_sizes) |
|
|
| |
| wrapper = Qwen3VisionForExport(vision, grid_thw) |
| wrapper = wrapper.to(dtype).eval().cuda() |
|
|
| |
| |
| pv_shape = captured_vit.pixel_values_shape |
| if batch_size > 1: |
| pv_shape = (pv_shape[0] * batch_size, pv_shape[1]) |
| pixel_values = torch.randn(pv_shape, dtype=dtype, device="cuda") |
|
|
| logger.info(f" pixel_values: {pixel_values.shape} ({pixel_values.dtype})") |
|
|
| precision_tag = "bf16" if use_bf16 else "fp32" |
| output_path = os.path.join(output_dir, f"vit_{precision_tag}.onnx") |
| os.makedirs(os.path.dirname(output_path), exist_ok=True) |
|
|
| output_names = ["image_embeds", "deepstack_features"] |
|
|
| logger.info(f" Exporting to {output_path}...") |
| with torch.inference_mode(): |
| torch.onnx.export( |
| wrapper, |
| (pixel_values,), |
| output_path, |
| input_names=["pixel_values"], |
| output_names=output_names, |
| opset_version=19, |
| do_constant_folding=True, |
| export_params=True, |
| dynamo=False, |
| ) |
|
|
| logger.info(" ViT exported successfully!") |
| _consolidate_external_data(output_path) |
| verify_onnx_export(output_path) |
|
|
| |
| |
| verify_onnx_with_ort( |
| onnx_path=output_path, |
| pytorch_module=wrapper, |
| sample_inputs={"pixel_values": pixel_values}, |
| output_names=output_names, |
| label="ViT", |
| ) |
|
|
| |
| _restore_vision_attention(vision, originals) |
|
|
| return output_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def export_llm_to_onnx(policy, captured_llm, output_dir, use_bf16=True, batch_size=1): |
| """Export the Qwen3-VL text model (LLM) to ONNX. |
| |
| The LLM receives inputs_embeds (with vision tokens already scattered in), |
| pre-computed 3D position_ids, and deepstack visual embeddings. Position ID |
| computation (get_rope_index) stays in PyTorch at runtime — only the |
| transformer layers are exported to ONNX/TRT. |
| |
| Deepstack injection is handled inside the wrapper: visual features are added |
| to hidden states at the first N layers (N = number of deepstack features, |
| typically 3) at positions indicated by visual_pos_masks. |
| |
| Input: |
| inputs_embeds: [B, seq_len, hidden_size] |
| attention_mask: [B, seq_len] (int64, 1=attend, 0=pad) |
| position_ids: [3, B, seq_len] (temporal, height, width for 3D RoPE) |
| visual_pos_masks: [B, seq_len] (bool, True at visual token positions) |
| deepstack_0: [num_vis_tokens, hidden_size] (deepstack feature for layer 0) |
| deepstack_1: [num_vis_tokens, hidden_size] (deepstack feature for layer 1) |
| deepstack_2: [num_vis_tokens, hidden_size] (deepstack feature for layer 2) |
| Output: |
| embeddings: [B, seq_len, hidden_size] |
| """ |
| logger.info("\n" + "=" * 80) |
| logger.info("Exporting LLM (Qwen3-VL Text Model) to ONNX") |
| logger.info("=" * 80) |
|
|
| backbone = policy.model.backbone |
| qwen_model = backbone.model |
| inner_model = qwen_model.model |
| text_model = inner_model.language_model |
| select_layer = backbone.select_layer |
|
|
| |
| from transformers.models.qwen3_vl.modeling_qwen3_vl import Qwen3VLTextRotaryEmbedding |
|
|
| text_config = copy.deepcopy(text_model.config) |
| text_config._attn_implementation = "eager" |
| text_config.num_hidden_layers = select_layer |
|
|
| logger.info( |
| f" LLM config: hidden_size={text_config.hidden_size}, " |
| f"num_layers={text_config.num_hidden_layers}, " |
| f"attn_implementation={text_config._attn_implementation}" |
| ) |
|
|
| |
| num_deepstack = 0 |
| if captured_llm.deepstack_visual_embeds is not None: |
| num_deepstack = len(captured_llm.deepstack_visual_embeds) |
| logger.info(f" Deepstack layers: {num_deepstack}") |
|
|
| class LLMForExport(torch.nn.Module): |
| """ONNX-exportable wrapper for Qwen3-VL text model with deepstack. |
| |
| Key adaptations from the original Qwen3VLTextModel: |
| 1. Eager attention (no flash) — avoids ONNX-incompatible flash attention |
| 2. Simple causal mask — avoids COMPLEX128 ops from HuggingFace's mask |
| 3. Deepstack injection via torch.where — avoids boolean indexing |
| 4. Position IDs as explicit input — get_rope_index() stays in PyTorch |
| 5. Deepstack features as separate tensor inputs (not a Python list) |
| """ |
|
|
| def __init__(self, config, n_deepstack): |
| super().__init__() |
| |
| self.layers = nn.ModuleList( |
| [ |
| |
| __import__( |
| "transformers.models.qwen3_vl.modeling_qwen3_vl", |
| fromlist=["Qwen3VLTextDecoderLayer"], |
| ).Qwen3VLTextDecoderLayer(config, layer_idx) |
| for layer_idx in range(config.num_hidden_layers) |
| ] |
| ) |
| |
| |
| |
| self.rotary_emb = Qwen3VLTextRotaryEmbedding(config=config) |
| self.n_deepstack = n_deepstack |
|
|
| def _simple_causal_mask(self, dtype, device, batch_size, seq_len, attention_mask): |
| """ONNX-compatible causal mask without complex type casts.""" |
| mask_value = torch.finfo(dtype).min * 0.5 |
| causal_mask = torch.triu( |
| torch.full((seq_len, seq_len), mask_value, device=device, dtype=dtype), |
| diagonal=1, |
| ) |
| causal_mask = causal_mask.unsqueeze(0).unsqueeze(0).expand(batch_size, 1, -1, -1) |
|
|
| if attention_mask is not None and attention_mask.dim() == 2: |
| padding_mask = attention_mask[:, None, None, :].to(dtype) |
| padding_mask = (1.0 - padding_mask) * mask_value |
| causal_mask = causal_mask + padding_mask |
|
|
| return causal_mask |
|
|
| def _deepstack_add(self, hidden_states, visual_pos_masks, visual_embeds): |
| """ONNX-friendly deepstack injection using torch.where. |
| |
| Original uses boolean indexing: |
| hidden_states[visual_pos_masks, :] += visual_embeds |
| This is not ONNX-friendly. Instead we: |
| 1. Build a full-size delta tensor (zeros except at visual positions) |
| 2. Add it to hidden_states |
| """ |
| |
| |
| |
| B, S, H = hidden_states.shape |
|
|
| |
| mask_expanded = visual_pos_masks.unsqueeze(-1) |
|
|
| |
| |
| delta = torch.zeros_like(hidden_states) |
| |
| delta = delta.masked_scatter(mask_expanded.expand_as(delta), visual_embeds) |
|
|
| hidden_states = hidden_states + delta |
| return hidden_states |
|
|
| def forward( |
| self, |
| inputs_embeds, |
| attention_mask, |
| position_ids, |
| visual_pos_masks=None, |
| deepstack_0=None, |
| deepstack_1=None, |
| deepstack_2=None, |
| ): |
| batch_size, seq_len = inputs_embeds.shape[:2] |
| dtype = inputs_embeds.dtype |
| device = inputs_embeds.device |
|
|
| |
| attn_mask = self._simple_causal_mask(dtype, device, batch_size, seq_len, attention_mask) |
|
|
| |
| text_position_ids = position_ids[0] |
|
|
| |
| cache_position = torch.arange(seq_len, device=device) |
|
|
| hidden_states = inputs_embeds |
|
|
| |
| position_embeddings = self.rotary_emb(hidden_states, position_ids) |
|
|
| |
| deepstack_list = [] |
| if deepstack_0 is not None: |
| deepstack_list.append(deepstack_0) |
| if deepstack_1 is not None: |
| deepstack_list.append(deepstack_1) |
| if deepstack_2 is not None: |
| deepstack_list.append(deepstack_2) |
|
|
| |
| for layer_idx, decoder_layer in enumerate(self.layers): |
| layer_outputs = decoder_layer( |
| hidden_states, |
| attention_mask=attn_mask, |
| position_ids=text_position_ids, |
| past_key_values=None, |
| cache_position=cache_position, |
| position_embeddings=position_embeddings, |
| ) |
| hidden_states = layer_outputs |
|
|
| |
| if visual_pos_masks is not None and layer_idx < len(deepstack_list): |
| hidden_states = self._deepstack_add( |
| hidden_states, visual_pos_masks, deepstack_list[layer_idx] |
| ) |
|
|
| |
| |
| return hidden_states |
|
|
| |
| wrapper = LLMForExport(text_config, num_deepstack) |
|
|
| |
| |
| |
| src_state = text_model.state_dict() |
| dst_state = wrapper.state_dict() |
| loaded, skipped = 0, 0 |
| for key in dst_state: |
| if key in src_state: |
| dst_state[key] = src_state[key] |
| loaded += 1 |
| else: |
| skipped += 1 |
| logger.warning(f" Key not found in source: {key}") |
| wrapper.load_state_dict(dst_state) |
| logger.info(f" Loaded {loaded} weight tensors, skipped {skipped}") |
|
|
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
| wrapper = wrapper.to(dtype).eval().cuda() |
|
|
| |
| seq_len = captured_llm.inputs_embeds.shape[1] |
| hidden_size = text_config.hidden_size |
|
|
| inputs_embeds = torch.randn(batch_size, seq_len, hidden_size, dtype=dtype, device="cuda") |
| attention_mask = torch.ones(batch_size, seq_len, dtype=torch.int64, device="cuda") |
| position_ids = torch.zeros(3, batch_size, seq_len, dtype=torch.int64, device="cuda") |
|
|
| export_inputs = [inputs_embeds, attention_mask, position_ids] |
| input_names = ["inputs_embeds", "attention_mask", "position_ids"] |
| |
| llm_dynamic_axes = { |
| "inputs_embeds": {1: "seq_len"}, |
| "attention_mask": {1: "seq_len"}, |
| "position_ids": {2: "seq_len"}, |
| "embeddings": {1: "seq_len"}, |
| } |
|
|
| if num_deepstack > 0 and captured_llm.visual_pos_masks is not None: |
| |
| vis_mask = captured_llm.visual_pos_masks.to(device="cuda") |
| if vis_mask.shape[0] == 1 and batch_size > 1: |
| vis_mask = vis_mask.expand(batch_size, -1) |
| export_inputs.append(vis_mask) |
| input_names.append("visual_pos_masks") |
| llm_dynamic_axes["visual_pos_masks"] = {1: "seq_len"} |
|
|
| for i in range(num_deepstack): |
| ds = captured_llm.deepstack_visual_embeds[i] |
| |
| |
| if batch_size > 1: |
| ds_dummy = torch.randn( |
| batch_size * ds.shape[0], ds.shape[1], dtype=dtype, device="cuda" |
| ) |
| else: |
| ds_dummy = torch.randn_like(ds, dtype=dtype, device="cuda") |
| export_inputs.append(ds_dummy) |
| name = f"deepstack_{i}" |
| input_names.append(name) |
|
|
| logger.info(" Export input shapes:") |
| for name, tensor in zip(input_names, export_inputs): |
| logger.info(f" {name}: {tensor.shape} ({tensor.dtype})") |
|
|
| precision_tag = "bf16" if use_bf16 else "fp32" |
| output_path = os.path.join(output_dir, f"llm_{precision_tag}.onnx") |
| os.makedirs(os.path.dirname(output_path), exist_ok=True) |
|
|
| logger.info(f" Exporting to {output_path}...") |
| with torch.inference_mode(): |
| torch.onnx.export( |
| wrapper, |
| tuple(export_inputs), |
| output_path, |
| input_names=input_names, |
| output_names=["embeddings"], |
| opset_version=19, |
| do_constant_folding=True, |
| dynamic_axes=llm_dynamic_axes, |
| export_params=True, |
| dynamo=False, |
| ) |
|
|
| logger.info(" LLM exported successfully!") |
| _consolidate_external_data(output_path) |
| verify_onnx_export(output_path) |
| return output_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def prepare_observation(policy, dataset, traj_idx=0): |
| """Prepare a single observation for inference.""" |
| logger.info(f"\nPreparing observation from trajectory {traj_idx}...") |
|
|
| traj = dataset[traj_idx] |
| modality_configs = policy.get_modality_config() |
|
|
| data_point = extract_step_data( |
| traj, 0, modality_configs=modality_configs, embodiment_tag=policy.embodiment_tag |
| ) |
|
|
| observation = {} |
| for key, value in data_point.states.items(): |
| observation[f"state.{key}"] = value |
| for key, value in data_point.images.items(): |
| observation[f"video.{key}"] = np.array(value) |
| for key in modality_configs["language"].modality_keys: |
| observation[key] = data_point.text |
|
|
| parsed_obs = parse_observation_gr00t(observation, modality_configs) |
| logger.info(" Observation prepared") |
| return parsed_obs |
|
|
|
|
| |
| |
| |
|
|
|
|
| def export_vl_self_attention_to_onnx(policy, output_dir, vl_seq_len, use_bf16=True, batch_size=1): |
| """Export the vl_self_attention (SelfAttentionTransformer) to ONNX. |
| |
| This module sits between VLLN and the DiT, transforming backbone |
| embeddings. If the model has no vl_self_attention (nn.Identity), skip. |
| |
| Input: hidden_states [B, T, backbone_embedding_dim] (T is dynamic) |
| Output: hidden_states [B, T, backbone_embedding_dim] |
| """ |
| vl_sa = policy.model.action_head.vl_self_attention |
| if isinstance(vl_sa, nn.Identity): |
| logger.info(" vl_self_attention is Identity — skipping export") |
| return None |
|
|
| logger.info("\n" + "=" * 80) |
| logger.info("Exporting VL Self-Attention to ONNX") |
| logger.info("=" * 80) |
|
|
| config = policy.model.action_head.config |
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
| model = vl_sa.to(dtype).eval().cuda() |
|
|
| hidden_states = torch.randn( |
| batch_size, vl_seq_len, config.backbone_embedding_dim, dtype=dtype, device="cuda" |
| ) |
| logger.info(f" hidden_states: {hidden_states.shape} ({hidden_states.dtype})") |
|
|
| output_path = os.path.join(output_dir, "vl_self_attention.onnx") |
| os.makedirs(os.path.dirname(output_path), exist_ok=True) |
|
|
| logger.info(f" Exporting to {output_path}...") |
| with torch.inference_mode(): |
| torch.onnx.export( |
| model, |
| (hidden_states,), |
| output_path, |
| input_names=["hidden_states"], |
| output_names=["output"], |
| dynamic_axes={ |
| "hidden_states": {1: "seq_len"}, |
| "output": {1: "seq_len"}, |
| }, |
| opset_version=19, |
| do_constant_folding=True, |
| dynamo=False, |
| ) |
|
|
| logger.info(" VL Self-Attention exported successfully!") |
| verify_onnx_export(output_path) |
| return output_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def export_state_encoder_to_onnx(policy, output_dir, use_bf16=True, batch_size=1): |
| """Export the state encoder (CategorySpecificMLP) to ONNX. |
| |
| N1.7 change: input_dim = max_state_dim * state_history_length. |
| The state is reshaped from [B, state_history_length, max_state_dim] |
| to [B, 1, state_history_length * max_state_dim] before encoding. |
| |
| Input: state [B, 1, max_state_dim * state_history_length], embodiment_id [B] |
| Output: [B, 1, input_embedding_dim] |
| """ |
| logger.info("\n" + "=" * 80) |
| logger.info("Exporting State Encoder to ONNX") |
| logger.info("=" * 80) |
|
|
| config = policy.model.action_head.config |
| state_encoder = policy.model.action_head.state_encoder |
|
|
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
| model = state_encoder.to(dtype).eval().cuda() |
|
|
| |
| state_input_dim = config.max_state_dim * config.state_history_length |
| state = torch.randn(batch_size, 1, state_input_dim, dtype=dtype, device="cuda") |
| embodiment_id = torch.zeros(batch_size, dtype=torch.int64, device="cuda") |
|
|
| logger.info(f" state: {state.shape} ({state.dtype})") |
| logger.info(f" embodiment_id: {embodiment_id.shape} ({embodiment_id.dtype})") |
| logger.info( |
| f" (max_state_dim={config.max_state_dim}, " |
| f"state_history_length={config.state_history_length})" |
| ) |
|
|
| output_path = os.path.join(output_dir, "state_encoder.onnx") |
| os.makedirs(os.path.dirname(output_path), exist_ok=True) |
|
|
| logger.info(f" Exporting to {output_path}...") |
| with torch.inference_mode(): |
| torch.onnx.export( |
| model, |
| (state, embodiment_id), |
| output_path, |
| input_names=["state", "embodiment_id"], |
| output_names=["output"], |
| opset_version=19, |
| do_constant_folding=True, |
| |
| |
| |
| |
| |
| dynamo=False, |
| ) |
|
|
| logger.info(" State Encoder exported successfully!") |
| verify_onnx_export(output_path) |
| return output_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def export_action_encoder_to_onnx(policy, output_dir, use_bf16=True, batch_size=1): |
| """Export the action encoder (MultiEmbodimentActionEncoder) to ONNX. |
| |
| Input: actions [B, action_horizon, max_action_dim], timesteps [B], embodiment_id [B] |
| Output: [B, action_horizon, input_embedding_dim] |
| """ |
| logger.info("\n" + "=" * 80) |
| logger.info("Exporting Action Encoder to ONNX") |
| logger.info("=" * 80) |
|
|
| config = policy.model.action_head.config |
| action_encoder = policy.model.action_head.action_encoder |
|
|
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
| model = action_encoder.to(dtype).eval().cuda() |
|
|
| actions = torch.randn( |
| batch_size, config.action_horizon, config.max_action_dim, dtype=dtype, device="cuda" |
| ) |
| timesteps = torch.zeros(batch_size, dtype=torch.int64, device="cuda") |
| embodiment_id = torch.zeros(batch_size, dtype=torch.int64, device="cuda") |
|
|
| logger.info(f" actions: {actions.shape} ({actions.dtype})") |
| logger.info(f" timesteps: {timesteps.shape} ({timesteps.dtype})") |
| logger.info(f" embodiment_id: {embodiment_id.shape} ({embodiment_id.dtype})") |
|
|
| output_path = os.path.join(output_dir, "action_encoder.onnx") |
| os.makedirs(os.path.dirname(output_path), exist_ok=True) |
|
|
| logger.info(f" Exporting to {output_path}...") |
| with torch.inference_mode(): |
| torch.onnx.export( |
| model, |
| (actions, timesteps, embodiment_id), |
| output_path, |
| input_names=["actions", "timesteps", "embodiment_id"], |
| output_names=["output"], |
| opset_version=19, |
| do_constant_folding=True, |
| dynamo=False, |
| ) |
|
|
| logger.info(" Action Encoder exported successfully!") |
| verify_onnx_export(output_path) |
| return output_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def export_dit_to_onnx(policy, captured_inputs, output_path, use_bf16=True, batch_size=1): |
| """Export the DiT (AlternateVLDiT) to ONNX. |
| |
| N1.7: image_mask and backbone_attention_mask are always present |
| (from Qwen3Backbone output). |
| |
| Input: sa_embs [B, sa_seq_len, input_embedding_dim], |
| vl_embs [B, vl_seq_len, backbone_embedding_dim], |
| timestep [B], image_mask [B, vl_seq_len], |
| backbone_attention_mask [B, vl_seq_len] |
| Output: [B, sa_seq_len, hidden_size] |
| """ |
| logger.info("\n" + "=" * 80) |
| logger.info("Exporting DiT to ONNX") |
| logger.info("=" * 80) |
|
|
| dit_model = policy.model.action_head.model |
| dit_model.eval() |
|
|
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
| dit_model = dit_model.to(dtype).cuda() |
|
|
| |
| sa_shape = (batch_size,) + captured_inputs.sa_embs.shape[1:] |
| vl_shape = (batch_size,) + captured_inputs.vl_embs.shape[1:] |
| ts_shape = (batch_size,) |
|
|
| sa_embs = torch.randn(sa_shape, dtype=dtype, device="cuda") |
| vl_embs = torch.randn(vl_shape, dtype=dtype, device="cuda") |
| timestep = torch.ones(ts_shape, dtype=torch.int64, device="cuda") |
|
|
| export_inputs = [sa_embs, vl_embs, timestep] |
| input_names = ["sa_embs", "vl_embs", "timestep"] |
| has_image_mask = captured_inputs.image_mask is not None |
| has_backbone_mask = captured_inputs.backbone_attention_mask is not None |
|
|
| if has_image_mask: |
| im_shape = (batch_size,) + captured_inputs.image_mask.shape[1:] |
| image_mask = torch.ones(im_shape, dtype=torch.bool, device="cuda") |
| export_inputs.append(image_mask) |
| input_names.append("image_mask") |
|
|
| if has_backbone_mask: |
| bm_shape = (batch_size,) + captured_inputs.backbone_attention_mask.shape[1:] |
| backbone_attention_mask = torch.ones(bm_shape, dtype=torch.bool, device="cuda") |
| export_inputs.append(backbone_attention_mask) |
| input_names.append("backbone_attention_mask") |
|
|
| logger.info(" Export input shapes:") |
| for name, tensor in zip(input_names, export_inputs): |
| logger.info(f" {name}: {tensor.shape} ({tensor.dtype})") |
|
|
| os.makedirs(os.path.dirname(output_path), exist_ok=True) |
|
|
| |
| |
| logger.info(f"Exporting to {output_path} with legacy ONNX exporter...") |
|
|
| |
| |
| |
| |
| |
| class DiTWrapper(torch.nn.Module): |
| def __init__(self, dit, use_image_mask, use_backbone_mask): |
| super().__init__() |
| self.dit = dit |
| self.use_image_mask = use_image_mask |
| self.use_backbone_mask = use_backbone_mask |
|
|
| def forward( |
| self, sa_embs, vl_embs, timestep, image_mask=None, backbone_attention_mask=None |
| ): |
| kwargs = {} |
| if self.use_image_mask and image_mask is not None: |
| kwargs["image_mask"] = image_mask |
| if self.use_backbone_mask and backbone_attention_mask is not None: |
| kwargs["backbone_attention_mask"] = backbone_attention_mask |
| return self.dit(sa_embs, vl_embs, timestep, **kwargs) |
|
|
| wrapped_model = DiTWrapper(dit_model, has_image_mask, has_backbone_mask) |
| wrapped_model.eval() |
|
|
| |
| |
| dit_dynamic_axes = { |
| "vl_embs": {1: "vl_seq_len"}, |
| } |
| if has_image_mask: |
| dit_dynamic_axes["image_mask"] = {1: "vl_seq_len"} |
| if has_backbone_mask: |
| dit_dynamic_axes["backbone_attention_mask"] = {1: "vl_seq_len"} |
|
|
| logger.info(f" Exporting to {output_path}...") |
| with torch.inference_mode(): |
| torch.onnx.export( |
| wrapped_model, |
| tuple(export_inputs), |
| output_path, |
| input_names=input_names, |
| output_names=["output"], |
| opset_version=19, |
| do_constant_folding=True, |
| export_params=True, |
| dynamic_axes=dit_dynamic_axes, |
| dynamo=False, |
| ) |
|
|
| logger.info(" DiT exported successfully!") |
|
|
| |
| |
| |
| _consolidate_external_data(output_path) |
|
|
| verify_onnx_export(output_path) |
| return output_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def export_action_decoder_to_onnx(policy, output_dir, use_bf16=True, batch_size=1): |
| """Export the action decoder (CategorySpecificMLP) to ONNX. |
| |
| Input: model_output [B, sa_seq_len, hidden_size], embodiment_id [B] |
| Output: [B, sa_seq_len, max_action_dim] |
| """ |
| logger.info("\n" + "=" * 80) |
| logger.info("Exporting Action Decoder to ONNX") |
| logger.info("=" * 80) |
|
|
| config = policy.model.action_head.config |
| action_decoder = policy.model.action_head.action_decoder |
|
|
| dtype = torch.bfloat16 if use_bf16 else torch.float32 |
| model = action_decoder.to(dtype).eval().cuda() |
|
|
| |
| sa_seq_len = 1 + config.action_horizon |
| model_output = torch.randn( |
| batch_size, sa_seq_len, config.hidden_size, dtype=dtype, device="cuda" |
| ) |
| embodiment_id = torch.zeros(batch_size, dtype=torch.int64, device="cuda") |
|
|
| logger.info(f" model_output: {model_output.shape} ({model_output.dtype})") |
| logger.info(f" embodiment_id: {embodiment_id.shape} ({embodiment_id.dtype})") |
|
|
| output_path = os.path.join(output_dir, "action_decoder.onnx") |
| os.makedirs(os.path.dirname(output_path), exist_ok=True) |
|
|
| logger.info(f" Exporting to {output_path}...") |
| with torch.inference_mode(): |
| torch.onnx.export( |
| model, |
| (model_output, embodiment_id), |
| output_path, |
| input_names=["model_output", "embodiment_id"], |
| output_names=["output"], |
| opset_version=19, |
| do_constant_folding=True, |
| dynamo=False, |
| ) |
|
|
| logger.info(" Action Decoder exported successfully!") |
| verify_onnx_export(output_path) |
| return output_path |
|
|
|
|
| |
| |
| |
|
|
|
|
| def main(args): |
| args.embodiment_tag = EmbodimentTag.resolve(args.embodiment_tag) |
| logger.info("=" * 80) |
| logger.info("GR00T N1.7 ONNX Export Script") |
| logger.info("=" * 80) |
| logger.info(f"Model path: {args.model_path}") |
| logger.info(f"Dataset path: {args.dataset_path}") |
| logger.info(f"Embodiment: {args.embodiment_tag}") |
| logger.info(f"Export mode: {args.export_mode}") |
| logger.info(f"Batch size: {args.batch_size}") |
| logger.info(f"Output directory: {args.output_dir}") |
| logger.info("=" * 80) |
|
|
| |
| logger.info("\n[Step 1] Loading policy...") |
| policy = Gr00tPolicy( |
| embodiment_tag=args.embodiment_tag, |
| model_path=args.model_path, |
| device="cuda", |
| ) |
| logger.info(" Policy loaded") |
|
|
| |
| logger.info("\n[Step 2] Loading dataset...") |
| dataset = LeRobotEpisodeLoader( |
| dataset_path=args.dataset_path, |
| modality_configs=policy.get_modality_config(), |
| ) |
| logger.info(f" Dataset loaded ({len(dataset)} trajectories)") |
|
|
| |
| logger.info("\n[Step 3] Capturing model inputs from actual inference...") |
|
|
| dit_capture = DiTInputCapture() |
| dit_hook = policy.model.action_head.model.register_forward_pre_hook( |
| dit_capture.hook_fn, with_kwargs=True |
| ) |
|
|
| |
| vit_capture = None |
| vit_hook = None |
| llm_capture = None |
| llm_hook = None |
| if args.export_mode == "full_pipeline": |
| vit_capture = ViTInputCapture() |
| qwen_model = policy.model.backbone.model |
| vit_hook = qwen_model.model.visual.register_forward_hook( |
| vit_capture.hook_fn, with_kwargs=True |
| ) |
|
|
| llm_capture = LLMInputCapture() |
| llm_hook = qwen_model.model.language_model.register_forward_pre_hook( |
| llm_capture.hook_fn, with_kwargs=True |
| ) |
|
|
| observation = prepare_observation(policy, dataset, traj_idx=0) |
| logger.info(" Running inference to capture shapes...") |
| with torch.inference_mode(): |
| _ = policy.get_action(observation) |
|
|
| dit_hook.remove() |
| if vit_hook is not None: |
| vit_hook.remove() |
| if llm_hook is not None: |
| llm_hook.remove() |
|
|
| if not dit_capture.captured: |
| logger.error(" Failed to capture DiT inputs!") |
| return |
| if args.export_mode == "full_pipeline" and not vit_capture.captured: |
| logger.error(" Failed to capture ViT inputs!") |
| return |
| if args.export_mode == "full_pipeline" and not llm_capture.captured: |
| logger.error(" Failed to capture LLM inputs!") |
| return |
|
|
| |
| action_head_config = policy.model.action_head.config |
| sa_seq_len = 1 + action_head_config.action_horizon |
| vl_seq_len = dit_capture.vl_embs.shape[1] |
|
|
| |
| num_patches = vit_capture.pixel_values_shape[0] if vit_capture and vit_capture.captured else 256 |
| num_merged_patches = vit_capture.output_shape[0] if vit_capture and vit_capture.captured else 64 |
| |
| llm_seq_len = ( |
| llm_capture.inputs_embeds.shape[1] if llm_capture and llm_capture.captured else vl_seq_len |
| ) |
| llm_hidden_size = ( |
| llm_capture.inputs_embeds.shape[2] if llm_capture and llm_capture.captured else 0 |
| ) |
| num_deepstack = ( |
| len(llm_capture.deepstack_visual_embeds) |
| if (llm_capture and llm_capture.deepstack_visual_embeds) |
| else 0 |
| ) |
| num_vis_tokens = llm_capture.deepstack_visual_embeds[0].shape[0] if num_deepstack > 0 else 0 |
|
|
| export_metadata = { |
| "schema_version": EXPORT_METADATA_SCHEMA_VERSION, |
| "model_version": "n1d7", |
| "sa_seq_len": int(sa_seq_len), |
| "vl_seq_len": int(vl_seq_len), |
| "llm_seq_len": int(llm_seq_len), |
| "llm_hidden_size": int(llm_hidden_size), |
| "num_deepstack": int(num_deepstack), |
| "num_vis_tokens": int(num_vis_tokens), |
| "num_patches": int(num_patches), |
| "num_merged_patches": int(num_merged_patches), |
| "action_horizon": int(action_head_config.action_horizon), |
| "max_action_dim": int(action_head_config.max_action_dim), |
| "max_state_dim": int(action_head_config.max_state_dim), |
| "state_history_length": int(action_head_config.state_history_length), |
| "hidden_size": int(action_head_config.hidden_size), |
| "input_embedding_dim": int(action_head_config.input_embedding_dim), |
| "backbone_embedding_dim": int(action_head_config.backbone_embedding_dim), |
| "embodiment_tag": str(args.embodiment_tag), |
| "export_mode": args.export_mode, |
| "precision": args.precision, |
| "batch_size": args.batch_size, |
| } |
| |
| |
| |
| if vit_capture is not None and vit_capture.captured: |
| export_metadata["vit_grid_thw"] = [ |
| [int(x) for x in row] for row in vit_capture.grid_thw.tolist() |
| ] |
| os.makedirs(args.output_dir, exist_ok=True) |
| metadata_path = os.path.join(args.output_dir, "export_metadata.json") |
| with open(metadata_path, "w") as f: |
| json.dump(export_metadata, f, indent=2) |
| logger.info(f" Saved export metadata to {metadata_path}") |
|
|
| |
| bs = args.batch_size |
| if args.export_mode == "dit_only": |
| logger.info("\n[Step 4] Exporting DiT to ONNX (dit_only mode)...") |
| dit_output_path = os.path.join(args.output_dir, "dit_bf16.onnx") |
| export_dit_to_onnx( |
| policy=policy, |
| captured_inputs=dit_capture, |
| output_path=dit_output_path, |
| use_bf16=True, |
| batch_size=bs, |
| ) |
|
|
| elif args.export_mode == "action_head": |
| logger.info("\n[Step 4] Exporting action head components to ONNX...") |
|
|
| |
| logger.info("\n--- [4a] State Encoder ---") |
| export_state_encoder_to_onnx(policy, args.output_dir, use_bf16=True, batch_size=bs) |
|
|
| |
| logger.info("\n--- [4b] Action Encoder ---") |
| export_action_encoder_to_onnx(policy, args.output_dir, use_bf16=True, batch_size=bs) |
|
|
| |
| logger.info("\n--- [4c] DiT ---") |
| dit_output_path = os.path.join(args.output_dir, "dit_bf16.onnx") |
| export_dit_to_onnx( |
| policy=policy, |
| captured_inputs=dit_capture, |
| output_path=dit_output_path, |
| use_bf16=True, |
| batch_size=bs, |
| ) |
|
|
| |
| logger.info("\n--- [4d] Action Decoder ---") |
| export_action_decoder_to_onnx(policy, args.output_dir, use_bf16=True, batch_size=bs) |
|
|
| elif args.export_mode == "full_pipeline": |
| logger.info("\n[Step 4] Exporting full pipeline to ONNX...") |
| logger.info(" (ViT TRT + LLM TRT + Action Head TRT)") |
|
|
| |
| |
| logger.info("\n--- [4a] ViT (Qwen3-VL Vision, FP32 for TRT accuracy) ---") |
| export_vit_to_onnx(policy, args.output_dir, vit_capture, use_bf16=False, batch_size=bs) |
|
|
| |
| logger.info("\n--- [4b] LLM (Qwen3-VL Text Model) ---") |
| export_llm_to_onnx(policy, llm_capture, args.output_dir, use_bf16=True, batch_size=bs) |
|
|
| |
| logger.info("\n--- [4c] VL Self-Attention ---") |
| export_vl_self_attention_to_onnx( |
| policy, args.output_dir, vl_seq_len=vl_seq_len, use_bf16=True, batch_size=bs |
| ) |
|
|
| |
| logger.info("\n--- [4d] State Encoder ---") |
| export_state_encoder_to_onnx(policy, args.output_dir, use_bf16=True, batch_size=bs) |
|
|
| |
| logger.info("\n--- [4e] Action Encoder ---") |
| export_action_encoder_to_onnx(policy, args.output_dir, use_bf16=True, batch_size=bs) |
|
|
| |
| logger.info("\n--- [4f] DiT ---") |
| dit_output_path = os.path.join(args.output_dir, "dit_bf16.onnx") |
| export_dit_to_onnx( |
| policy=policy, |
| captured_inputs=dit_capture, |
| output_path=dit_output_path, |
| use_bf16=True, |
| batch_size=bs, |
| ) |
|
|
| |
| logger.info("\n--- [4g] Action Decoder ---") |
| export_action_decoder_to_onnx(policy, args.output_dir, use_bf16=True, batch_size=bs) |
|
|
| |
| logger.info("\n" + "=" * 80) |
| logger.info("EXPORT COMPLETE!") |
| logger.info("=" * 80) |
| logger.info(f"\nExported files in: {args.output_dir}") |
|
|
| for f in sorted(os.listdir(args.output_dir)): |
| fpath = os.path.join(args.output_dir, f) |
| if os.path.isfile(fpath): |
| size_mb = os.path.getsize(fpath) / (1024 * 1024) |
| logger.info(f" {f}: {size_mb:.2f} MB") |
|
|
|
|
| @dataclass |
| class ExportConfig: |
| """Configuration for exporting GR00T N1.7 model to ONNX.""" |
|
|
| model_path: str |
| """Path to the model checkpoint (required).""" |
|
|
| dataset_path: str |
| """Path to the dataset (required, used to capture input shapes).""" |
|
|
| embodiment_tag: Optional[EmbodimentTag] = None |
| """Embodiment tag. If not provided, auto-detected from model's processor_config.json.""" |
|
|
| output_dir: str = "./gr00t_trt_deployment/onnx" |
| """Output directory for ONNX models.""" |
|
|
| export_mode: ExportMode = ExportMode.dit_only |
| """Export mode: 'dit_only', 'action_head' (4 components), or 'full_pipeline' (ViT + action head).""" |
|
|
| precision: Literal["bf16"] = "bf16" |
| """Export precision for the generated ONNX graph. |
| |
| Currently fixed to 'bf16': every Step 4 exporter passes a hardcoded |
| `use_bf16=` argument and does not read this field beyond writing it |
| into export_metadata.json. Re-introducing 'fp16'/'fp32'/'fp8' here |
| requires plumbing this field into each exporter first; until then |
| the Literal is narrowed so the CLI cannot accept a value the export |
| will ignore.""" |
|
|
| batch_size: int = 1 |
| """Batch size baked into the exported ONNX models (default: 1).""" |
|
|
|
|
| if __name__ == "__main__": |
| args = tyro.cli(ExportConfig) |
| if args.embodiment_tag is None: |
| |
| config_file = Path(args.model_path) / "processor_config.json" |
| if not config_file.exists(): |
| raise ValueError( |
| f"Cannot auto-detect embodiment_tag: {config_file} not found. " |
| "Please provide --embodiment-tag explicitly." |
| ) |
| with open(config_file, "r") as f: |
| processor_config = json.load(f) |
| modality_configs = processor_config.get("processor_kwargs", {}).get("modality_configs", {}) |
| if len(modality_configs) == 0: |
| raise ValueError( |
| "Cannot auto-detect embodiment_tag: no modality_configs found in processor_config.json. " |
| "Please provide --embodiment-tag explicitly." |
| ) |
| if len(modality_configs) == 1: |
| embodiment_key = next(iter(modality_configs)) |
| args.embodiment_tag = EmbodimentTag.resolve(embodiment_key) |
| logger.info( |
| f"Auto-detected embodiment tag: {args.embodiment_tag} (from {embodiment_key})" |
| ) |
| else: |
| available = sorted(modality_configs.keys()) |
| raise ValueError( |
| f"Multiple embodiments found in processor_config.json: {available}. " |
| "Please provide --embodiment-tag explicitly." |
| ) |
| main(args) |
|
|