From 062ea72dffd7e21720f0657d8f4788e33a524e89 Mon Sep 17 00:00:00 2001 From: liyh15 <73781551+xgbah@users.noreply.github.com> Date: Fri, 9 Oct 2026 09:32:34 +0800 Subject: [PATCH 1/4] feat(openvla): add OpenVLA optimization catalog and implementations --- .../test_openvla_skip_fa2_unpad.py | 368 ++++++++++ .../optimizations/test_openvla_vision_timm.py | 205 ++++++ .../optimizations/models/openvla/__init__.py | 8 + .../optimizations/models/openvla/catalog.py | 200 ++++++ .../models/openvla/compile_fsdp1.py | 377 ++++++++++ .../models/openvla/fixbf16support.py | 36 + .../models/openvla/fsdp_prefetch.py | 88 +++ .../models/openvla/fusedAdamW.py | 15 + .../optimizations/models/openvla/gc_freeze.py | 51 ++ .../models/openvla/reproducibility.py | 673 ++++++++++++++++++ .../models/openvla/skip_fa2_unpad.py | 59 ++ .../models/openvla/spawn_dataloader.py | 270 +++++++ .../models/openvla/text_len_bucket.py | 107 +++ .../models/openvla/vision_timm.py | 98 +++ 14 files changed, 2555 insertions(+) create mode 100644 test/optimizations/test_openvla_skip_fa2_unpad.py create mode 100644 test/optimizations/test_openvla_vision_timm.py create mode 100644 turbo_physai/optimizations/models/openvla/__init__.py create mode 100644 turbo_physai/optimizations/models/openvla/catalog.py create mode 100644 turbo_physai/optimizations/models/openvla/compile_fsdp1.py create mode 100644 turbo_physai/optimizations/models/openvla/fixbf16support.py create mode 100644 turbo_physai/optimizations/models/openvla/fsdp_prefetch.py create mode 100644 turbo_physai/optimizations/models/openvla/fusedAdamW.py create mode 100644 turbo_physai/optimizations/models/openvla/gc_freeze.py create mode 100644 turbo_physai/optimizations/models/openvla/reproducibility.py create mode 100644 turbo_physai/optimizations/models/openvla/skip_fa2_unpad.py create mode 100644 turbo_physai/optimizations/models/openvla/spawn_dataloader.py create mode 100644 turbo_physai/optimizations/models/openvla/text_len_bucket.py create mode 100644 turbo_physai/optimizations/models/openvla/vision_timm.py diff --git a/test/optimizations/test_openvla_skip_fa2_unpad.py b/test/optimizations/test_openvla_skip_fa2_unpad.py new file mode 100644 index 0000000..d91d41b --- /dev/null +++ b/test/optimizations/test_openvla_skip_fa2_unpad.py @@ -0,0 +1,368 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""Tests for the ``openvla.llm.skip_fa2_unpad`` Group. + +The Group installs ``make_fast_fa2_causal_mask_wrapper`` in place of +``LlamaModel._update_causal_mask``. On a right-padded prefill, transformers +hands the 2D padding mask straight back on its FA2 branch, which pushes +``LlamaFlashAttention2._flash_attention_forward`` onto ``_get_unpad_data`` + +``flash_attn_varlen_func`` + ``pad_input``. The replacement returns ``None`` +instead, so the same call takes plain ``flash_attn_func(causal=True)``. + +Three separate properties carry that behaviour, and each needs its own test +because none of them can observe the others: + +* **It is installed and it fires.** A test that only compares the two FA2 + arms would also pass with the Group disabled, so the padding case asserts + the unpad path is really taken without the replacement and really skipped + with it. +* **The dropped mask is equivalent.** Right padding under a causal mask + confines every real token to real tokens, so real-token outputs, the loss + and every parameter gradient must match bit for bit. +* **Every other path is delegated.** A wrapper that also returned ``None`` + while decoding, or under ``sdpa``/``eager``, would silently drop a real + padding mask. Those branches are pinned without a device by driving the + wrapper with a stub ``original``. +""" + +from __future__ import annotations + +import textwrap +import types + +import pytest + + +torch = pytest.importorskip("torch") + +from turbo_physai.engine.checking.context import detect_context +from turbo_physai.engine.checking.ordering import Preparation +from turbo_physai.engine.config.loader import load_optimization_config +from turbo_physai.engine.contracts import Decision, Mechanism +from turbo_physai.engine.definitions.registry import default_registry +from turbo_physai.engine.execution.replacements import default_handlers +from turbo_physai.optimizations.models.openvla.catalog import SKIP_FA2_UNPAD +from turbo_physai.optimizations.models.openvla.skip_fa2_unpad import ( + make_fast_fa2_causal_mask_wrapper, +) + + +GROUP_ID = "openvla.llm.skip_fa2_unpad" +TARGET = "transformers.models.llama.modeling_llama.LlamaModel._update_causal_mask" +REPLACEMENT = ( + "turbo_physai.optimizations.models.openvla." + "skip_fa2_unpad.make_fast_fa2_causal_mask_wrapper" +) + +SEED = 20240607 +LAYERS = 2 +HEADS = 4 +HEAD_DIM = 16 +HIDDEN = HEADS * HEAD_DIM +VOCAB = 128 + + +class _OriginalRecorder: + """Stub ``LlamaModel._update_causal_mask`` that records how it was called.""" + + def __init__(self): + self.calls = [] + self.result = object() + + def __call__( + self, model, attention_mask, input_tensor, cache_position, past_seen_tokens + ): + self.calls.append( + (model, attention_mask, input_tensor, cache_position, past_seen_tokens) + ) + return self.result + + +def _stub_model(implementation): + """A stand-in for ``self``: the wrapper only reads ``config``.""" + + return types.SimpleNamespace( + config=types.SimpleNamespace(_attn_implementation=implementation) + ) + + +def _padded_mask(lengths, sequence_length): + """2D right-padding mask: real tokens first, padding afterwards.""" + + mask = torch.zeros(len(lengths), sequence_length, dtype=torch.long) + for row, length in enumerate(lengths): + mask[row, :length] = 1 + return mask + + +def _left_padded_mask(lengths, sequence_length): + """2D left-padding mask: padding first, real tokens afterwards.""" + + mask = torch.zeros(len(lengths), sequence_length, dtype=torch.long) + for row, length in enumerate(lengths): + mask[row, sequence_length - length :] = 1 + return mask + + +@pytest.mark.parametrize( + "options", [None, {}, {"unused": "option"}], ids=["no-options", "empty", "present"] +) +def test_prefill_padding_mask_short_circuits_the_causal_mask(options): + """FA2 prefill with a mask returns ``None`` and never reaches ``original``.""" + + original = _OriginalRecorder() + wrapper = make_fast_fa2_causal_mask_wrapper(original, options) + + causal_mask = wrapper( + _stub_model("flash_attention_2"), + _padded_mask((4, 2), 4), + torch.zeros(2, 4, 8), + torch.arange(4), + 0, + ) + + assert causal_mask is None + assert original.calls == [], "the original must not materialise a mask" + + +@pytest.mark.parametrize( + "implementation,past_seen_tokens,mask_present", + [ + ("eager", 0, True), + ("sdpa", 0, True), + ("flash_attention_2", 4, True), + ("flash_attention_2", 0, False), + ], + ids=["eager", "sdpa", "cached-decode", "no-mask"], +) +def test_every_other_path_delegates_to_the_original( + implementation, past_seen_tokens, mask_present +): + """Non-prefill, non-FA2 and mask-less calls stay bit-for-bit untouched.""" + + original = _OriginalRecorder() + wrapper = make_fast_fa2_causal_mask_wrapper(original, None) + model = _stub_model(implementation) + mask = _padded_mask((4, 2), 4) if mask_present else None + input_tensor = torch.zeros(2, 1 if past_seen_tokens else 4, 8) + cache_position = torch.tensor([4]) if past_seen_tokens else torch.arange(4) + + result = wrapper(model, mask, input_tensor, cache_position, past_seen_tokens) + + assert result is original.result + (call,) = original.calls + assert call[0] is model + assert call[1] is mask + assert call[2] is input_tensor + assert call[3] is cache_position + assert call[4] == past_seen_tokens + + +def test_left_padded_prefill_delegates_to_the_original(): + """Left padding must fall back: dropping the mask would expose the pad slots. + + With left padding the real tokens follow the pad slots, so plain causal + attention would attend to the pad key/values instead of having them + unpadded away by the varlen path -- a silent numerical regression. + """ + + original = _OriginalRecorder() + wrapper = make_fast_fa2_causal_mask_wrapper(original, None) + model = _stub_model("flash_attention_2") + mask = _left_padded_mask((4, 2), 4) + + result = wrapper(model, mask, torch.zeros(2, 4, 8), torch.arange(4), 0) + + assert result is original.result + (call,) = original.calls + assert call[1] is mask + + +def test_unpadded_prefill_still_short_circuits_the_causal_mask(): + """A batch with no padding at all has no left padding, so it still skips.""" + + original = _OriginalRecorder() + wrapper = make_fast_fa2_causal_mask_wrapper(original, None) + + causal_mask = wrapper( + _stub_model("flash_attention_2"), + _padded_mask((4, 4), 4), + torch.zeros(2, 4, 8), + torch.arange(4), + 0, + ) + + assert causal_mask is None + assert original.calls == [] + + +def test_engine_prepares_the_group_against_the_real_llama_target(tmp_path): + """The declaration resolves to the real target through the public entry point.""" + + pytest.importorskip("transformers.models.llama.modeling_llama") + + (spec,) = SKIP_FA2_UNPAD.specs + assert SKIP_FA2_UNPAD.group_id == GROUP_ID + assert spec.mechanism is Mechanism.WRAPPER + assert spec.target == TARGET + assert spec.replacement == REPLACEMENT + + config = tmp_path / "optimization.yaml" + config.write_text( + textwrap.dedent( + f""" + schema_version: turbophysai/optimization-config/v1 + kind: OptimizationConfig + metadata: {{id: openvla-skip-fa2-unpad, version: "1"}} + optimization_groups: + - id: {GROUP_ID} + """ + ), + encoding="utf-8", + ) + + # Preparation resolves the target and constructs the wrapper without + # installing it. The public `check()` entry point was removed upstream + # (`refactor(api)!: remove redundant check and CLI commands`); `apply()` + # would install the Group process-wide, which this test must not do. + prepared = Preparation(default_registry, default_handlers()).prepare( + run_id="openvla-skip-fa2-unpad", + config=load_optimization_config(config), + environment=detect_context(), + ) + + (group,) = prepared.groups + assert group.group_id == GROUP_ID + assert group.decision is Decision.APPLY + assert group.members == (f"{GROUP_ID}.update_causal_mask",) + assert prepared.conflicts == () + + +def _flash_attention_llama(lengths=(8, 5, 3), sequence_length=8): + """A small bf16 FA2 ``LlamaModel`` plus one right-padded token batch.""" + + llama = pytest.importorskip("transformers.models.llama.modeling_llama") + transformers = pytest.importorskip("transformers") + if not torch.cuda.is_available(): + pytest.skip("the FA2 varlen kernel needs a real accelerator device") + + config = transformers.LlamaConfig( + vocab_size=VOCAB, + hidden_size=HIDDEN, + intermediate_size=2 * HIDDEN, + num_hidden_layers=LAYERS, + num_attention_heads=HEADS, + num_key_value_heads=HEADS // 2, + max_position_embeddings=64, + _attn_implementation="flash_attention_2", + ) + model = ( + transformers.LlamaModel(config) + .to(device="cuda", dtype=torch.bfloat16) + .eval() + ) + assert model.config._attn_implementation == "flash_attention_2", ( + "the comparison would measure a different attention kernel" + ) + + generator = torch.Generator(device="cuda").manual_seed(SEED) + ids = torch.randint( + VOCAB, (len(lengths), sequence_length), generator=generator, device="cuda" + ) + mask = _padded_mask(lengths, sequence_length).to("cuda") + return llama, model, ids, mask, mask.bool() + + +def _count_unpad_calls(monkeypatch, llama): + """Count ``_get_unpad_data`` calls; the varlen path makes one per layer.""" + + calls = [] + genuine = llama._get_unpad_data + + def spy(attention_mask): + calls.append(attention_mask) + return genuine(attention_mask) + + monkeypatch.setattr(llama, "_get_unpad_data", spy) + return calls + + +@pytest.mark.hcu +@pytest.mark.model_deps +def test_padded_prefill_skips_the_unpad_path_and_stays_bit_identical(monkeypatch): + """The Group removes the varlen call and changes no number that is trained on.""" + + llama, model, ids, mask, real = _flash_attention_llama() + unpad_calls = _count_unpad_calls(monkeypatch, llama) + original = llama.LlamaModel._update_causal_mask + replacement = make_fast_fa2_causal_mask_wrapper(original, None) + + def arm(wrapped): + monkeypatch.setattr( + llama.LlamaModel, + "_update_causal_mask", + replacement if wrapped else original, + ) + model.zero_grad(set_to_none=True) + unpad_calls.clear() + + output = model(input_ids=ids, attention_mask=mask).last_hidden_state + calls = len(unpad_calls) + loss = (output * real.unsqueeze(-1)).sum() + loss.backward() + + gradients = {} + for name, parameter in model.named_parameters(): + assert parameter.grad is not None, f"no gradient reached {name}" + gradients[name] = parameter.grad.detach().clone() + return output.detach(), loss.detach().clone(), gradients, calls + + default_output, default_loss, default_gradients, default_calls = arm(False) + wrapped_output, wrapped_loss, wrapped_gradients, wrapped_calls = arm(True) + + # Without both halves the test could pass by running one code path twice: + # the default arm must really enter varlen, and the Group must really skip it. + assert default_calls > 0, "the default arm never reached the unpad path" + assert wrapped_calls == 0, "the replacement did not skip the unpad path" + + assert torch.equal(default_output[real], wrapped_output[real]) + assert torch.equal(default_loss, wrapped_loss) + assert set(default_gradients) == set(wrapped_gradients) + assert any( + float(gradient.abs().sum()) != 0.0 for gradient in default_gradients.values() + ), "every gradient is zero, which would make the comparison vacuous" + for name, gradient in default_gradients.items(): + assert torch.equal(gradient, wrapped_gradients[name]), ( + f"parameter gradient differs: {name}" + ) + + +@pytest.mark.hcu +@pytest.mark.model_deps +def test_unpadded_batch_never_enters_the_unpad_path(monkeypatch): + """A fully real batch already skipped the mask, so the Group is a no-op.""" + + llama, model, ids, mask, _ = _flash_attention_llama() + unpad_calls = _count_unpad_calls(monkeypatch, llama) + original = llama.LlamaModel._update_causal_mask + replacement = make_fast_fa2_causal_mask_wrapper(original, None) + unpadded = torch.ones_like(mask) + + def arm(wrapped): + monkeypatch.setattr( + llama.LlamaModel, + "_update_causal_mask", + replacement if wrapped else original, + ) + unpad_calls.clear() + with torch.no_grad(): + output = model(input_ids=ids, attention_mask=unpadded).last_hidden_state + return output, len(unpad_calls) + + default_output, default_calls = arm(False) + wrapped_output, wrapped_calls = arm(True) + + assert default_calls == 0 + assert wrapped_calls == 0 + assert torch.equal(default_output, wrapped_output) diff --git a/test/optimizations/test_openvla_vision_timm.py b/test/optimizations/test_openvla_vision_timm.py new file mode 100644 index 0000000..a522191 --- /dev/null +++ b/test/optimizations/test_openvla_vision_timm.py @@ -0,0 +1,205 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""Tests for the ``openvla.compile.fsdp1`` vision-tower rewrite. + +Prismatic swaps each ViT tower's ``forward`` for +``partial(featurizer.get_intermediate_layers, n={len(blocks) - 2})``, and the +FSDP1 compile flow whole-model-``torch.compile``s that tower. On the first real +forward Dynamo traces timm's + + take_indices = set(range(num_blocks - n, num_blocks) if isinstance(n, int) else n) + +whose ``set(...)`` over an already-sourced set trips ``assert source is None`` in +``torch/_dynamo/variables/base.py``. ``vision_timm`` reimplements the method +with a list instead of a set, which Dynamo traces cleanly. + +``n`` is the only thing the rewrite changes, so the equivalence claim is exactly +"the same blocks, in the same order, with the same numbers, for every ``n`` form +the monkey-patch can pass". Two properties are invisible to a bit-identity +check and are therefore asserted separately: + +* **The swap must actually happen.** The rewrite is behaviour-equivalent, so a + run where nothing was installed compares equal to itself and passes. +* **With ``options.compile`` off, the factory must hand back the original + method.** The towers stay eager there, so timm must be left globally + untouched. An accidental install would still produce identical numbers, yet + would mutate ``timm.models.vision_transformer.VisionTransformer`` for the + whole process -- the one consequence a numerical comparison can never see. +""" + +from __future__ import annotations + +import inspect + +import pytest + + +torch = pytest.importorskip("torch") + +from turbo_physai.engine.contracts import Mechanism +from turbo_physai.optimizations.models.openvla.catalog import COMPILE_FSDP1 +from turbo_physai.optimizations.models.openvla.vision_timm import ( + dynamo_safe_intermediate_layers, + timm_intermediate_layers_wrapper, +) + + +# timm ships as part of the upstream model stack rather than the repository +# requirements, so only its absence skips. A broken install must fail loudly. +timm_vision_transformer = pytest.importorskip("timm.models.vision_transformer") +VisionTransformer = timm_vision_transformer.VisionTransformer + + +pytestmark = pytest.mark.model_deps + + +DEPTH = 4 +SEED = 20240607 +TARGET = "timm.models.vision_transformer.VisionTransformer._intermediate_layers" +REPLACEMENT = ( + "turbo_physai.optimizations.models.openvla." + "vision_timm.timm_intermediate_layers_wrapper" +) + +# Every form the monkey-patch can pass: OpenVLA's set, plus the int and generic +# iterables `timm_intermediate_layers_wrapper` has to tolerate. `int-over` and +# `int-zero` pin the two range boundaries the rewrite reproduces implicitly. +INDEX_FORMS = [ + (1, "int-1"), + (2, "int-2"), + (DEPTH, "int-all"), + (DEPTH + 2, "int-beyond-depth"), + (0, "int-zero"), + ({DEPTH - 2}, "set-openvla"), + ({1, 3}, "set-two"), + ([1, 3], "list"), + ((0, 2), "tuple"), + (range(0, DEPTH, 2), "range"), +] + + +def _verify_timm_method(): + """Fail loudly if timm changed, or if an earlier test left timm patched.""" + + source = inspect.getsource(VisionTransformer._intermediate_layers) + assert "take_indices = set(" in source, ( + "timm's _intermediate_layers no longer matches the implementation this " + "rewrite mirrors; re-derive vision_timm.py before trusting this comparison. " + "A swapped-in method here also means an earlier test left timm patched." + ) + + +def _timm_method(): + """timm's own method, verified to still be the ``set(...)`` version we mirror.""" + + _verify_timm_method() + return VisionTransformer._intermediate_layers + + +def _vision_tower(): + """A tiny timm ViT tower, as ``eval`` so no dropout can make arms incomparable.""" + + return VisionTransformer( + img_size=32, + patch_size=8, + in_chans=3, + embed_dim=32, + depth=DEPTH, + num_heads=2, + num_classes=0, + ).eval() + + +def _input(): + generator = torch.Generator().manual_seed(SEED) + return torch.randn(2, 3, 32, 32, generator=generator) + + +def _flatten(value): + """Flatten the nested tuples ``get_intermediate_layers`` can return.""" + + if isinstance(value, (list, tuple)): + flattened = [] + for item in value: + flattened.extend(_flatten(item)) + return flattened + return [value] + + +@pytest.mark.parametrize("n,case", INDEX_FORMS, ids=[case for _, case in INDEX_FORMS]) +def test_rewrite_returns_the_same_blocks_for_every_index_form(n, case, monkeypatch): + """``set(n)`` -> ``list(n)`` must not move, drop or reorder a single block.""" + + original = _timm_method() + model, x = _vision_tower(), _input() + + with torch.no_grad(): + expected = original(model, x, n) + monkeypatch.setattr( + VisionTransformer, "_intermediate_layers", dynamo_safe_intermediate_layers + ) + with torch.no_grad(): + actual = dynamo_safe_intermediate_layers(model, x, n) + + assert len(actual) == len(expected), f"different block count for {case}" + for index, (left, right) in enumerate(zip(expected, actual)): + assert left.shape == right.shape, f"block {index} changed shape for {case}" + assert torch.equal(left, right), f"block {index} differs for {case}" + + +@pytest.mark.parametrize( + "keywords", + [{}, {"norm": True}, {"return_prefix_tokens": True}, {"reshape": True}], + ids=["plain", "norm", "prefix-tokens", "reshape"], +) +def test_prismatic_call_path_is_bit_identical(keywords, monkeypatch): + """The path OpenVLA really calls post-processes the same blocks unchanged.""" + + _verify_timm_method() + model, x = _vision_tower(), _input() + + with torch.no_grad(): + expected = model.get_intermediate_layers(x, n={DEPTH - 2}, **keywords) + monkeypatch.setattr( + VisionTransformer, "_intermediate_layers", dynamo_safe_intermediate_layers + ) + with torch.no_grad(): + actual = model.get_intermediate_layers(x, n={DEPTH - 2}, **keywords) + + left, right = _flatten(expected), _flatten(actual) + assert len(left) == len(right) + for index, (a, b) in enumerate(zip(left, right)): + assert a.shape == b.shape, f"tensor {index} changed shape" + assert torch.equal(a, b), f"tensor {index} differs" + + +def test_factory_installs_the_rewrite_only_when_compile_is_enabled(): + """Eager runs must leave timm globally untouched; compiled ones must swap.""" + + original = _timm_method() + + assert timm_intermediate_layers_wrapper(original, None) is original + assert timm_intermediate_layers_wrapper(original, {}) is original + assert timm_intermediate_layers_wrapper(original, {"compile": False}) is original + assert ( + timm_intermediate_layers_wrapper(original, {"compile": True}) + is dynamo_safe_intermediate_layers + ) + # Building the replacement must not mutate timm itself: installing it is the + # engine's job, and a factory that also swapped the class attribute would + # patch every timm ViT in the process before the Group was even applied. + assert VisionTransformer._intermediate_layers is original + + +def test_group_declares_the_timm_rewrite(): + """The declaration must keep pointing at the real timm method.""" + + (spec,) = [ + item + for item in COMPILE_FSDP1.specs + if item.target.startswith("timm.models.vision_transformer.") + ] + assert spec.mechanism is Mechanism.WRAPPER + assert spec.target == TARGET + assert spec.replacement == REPLACEMENT diff --git a/turbo_physai/optimizations/models/openvla/__init__.py b/turbo_physai/optimizations/models/openvla/__init__.py new file mode 100644 index 0000000..66de2ff --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/__init__.py @@ -0,0 +1,8 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""OpenVLA model optimization catalog.""" + +from . import catalog + +__all__ = ["catalog"] diff --git a/turbo_physai/optimizations/models/openvla/catalog.py b/turbo_physai/optimizations/models/openvla/catalog.py new file mode 100644 index 0000000..5182ca6 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/catalog.py @@ -0,0 +1,200 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""openvla optimization declarations. + +Add model-specific Groups here. Importing this module registers the declarations +with TurboPhysAI; the generated OptimizationConfig loads it through optimization_modules. + +Declarations only describe target/replacement pairs as strings; nothing here imports +OpenVLA/prismatic and nothing resolves runtime objects at import time. +""" + +from __future__ import annotations + +from ....engine.definitions import group, replace, wrap + +# --- BF16 mixed-precision support detection --------------------------------- +BF16_SUPPORT = group( + "openvla.bf16_support", + replace( + target="prismatic.util.torch_utils.check_bloat16_supported", + aliases=( + "prismatic.util.check_bloat16_supported", + "prismatic.training.strategies.base_strategy.check_bloat16_supported", + ), + replacement="turbo_physai.optimizations.models.openvla.fixbf16support.check_bloat16_supported", + ), +) + +# --- FSDP1 per-layer torch.compile ------------------------------------------- +COMPILE_FSDP1 = group( + "openvla.compile.fsdp1", + wrap( + target="prismatic.training.strategies.fsdp.FSDPStrategy.run_setup", + replacement="turbo_physai.optimizations.models.openvla.compile_fsdp1.run_setup_wrapper", + ), + wrap( + target="prismatic.training.strategies.fsdp.FSDPStrategy.save_checkpoint", + replacement="turbo_physai.optimizations.models.openvla.compile_fsdp1.save_checkpoint_wrapper", + ), + wrap( + target="timm.models.vision_transformer.VisionTransformer._intermediate_layers", + replacement="turbo_physai.optimizations.models.openvla.vision_timm.timm_intermediate_layers_wrapper", + ), +) + +# --- Fused AdamW (Group enabled => fused is the default) -------------------- +FUSED_ADAMW = group( + "openvla.adamw.fused", + wrap( + target="torch.optim.AdamW", + replacement="turbo_physai.optimizations.models.openvla.fusedAdamW.adamw_fused_wrapper", + ), +) + +# --- FSDP1 communication overlap (limit_all_gathers / fwd+bwd prefetch) ------ +FSDP_PREFETCH = group( + "openvla.fsdp.prefetch", + wrap( + target="torch.distributed.fsdp.FullyShardedDataParallel", + aliases=("prismatic.training.strategies.fsdp.FSDP",), + replacement="turbo_physai.optimizations.models.openvla.fsdp_prefetch.fsdp_prefetch_wrapper", + ), +) + +# --- Fixed-length (bucketed) text padding ------------------------------------ +TEXT_LEN_BUCKET = group( + "openvla.data.text_len_bucket", + wrap( + target="prismatic.util.data_utils.PaddedCollatorForActionPrediction.__call__", + replacement="turbo_physai.optimizations.models.openvla.text_len_bucket.bucketed_collate_wrapper", + ), +) + +# --- Skip FA2 varlen (unpad) on right-padded prefill ------------------------- +SKIP_FA2_UNPAD = group( + "openvla.llm.skip_fa2_unpad", + wrap( + target="transformers.models.llama.modeling_llama.LlamaModel._update_causal_mask", + replacement="turbo_physai.optimizations.models.openvla.skip_fa2_unpad.make_fast_fa2_causal_mask_wrapper", + ), +) + +# --- RLDS DataLoader: spawned workers (Group options: num_workers, default 1) -- +# 仅对 RLDS 数据集(RLDSDataset / EpisodicRLDSDataset)的 DataLoader 强制 +# `num_workers=N` + `spawn` context(RLDS 的 TF graph 不能 fork,必须 spawn 重建); +# RLDS 数据集类包成可 pickle 子类(序列化构造参数、worker 里重建 TF graph)。 +DATALOADER_SPAWN = group( + "openvla.dataloader.spawn", + wrap( + target="prismatic.vla.datasets.datasets.RLDSDataset", + aliases=( + "prismatic.vla.datasets.RLDSDataset", + "prismatic.vla.materialize.RLDSDataset", + ), + replacement="turbo_physai.optimizations.models.openvla.spawn_dataloader.rlds_dataset_spawn_wrapper", + ), + wrap( + target="prismatic.vla.datasets.datasets.EpisodicRLDSDataset", + aliases=( + "prismatic.vla.datasets.EpisodicRLDSDataset", + "prismatic.vla.materialize.EpisodicRLDSDataset", + ), + replacement="turbo_physai.optimizations.models.openvla.spawn_dataloader.episodic_rlds_dataset_spawn_wrapper", + ), + wrap( + target="torch.utils.data.DataLoader", + aliases=("prismatic.training.strategies.base_strategy.DataLoader",), + replacement="turbo_physai.optimizations.models.openvla.spawn_dataloader.dataloader_spawn_wrapper", + ), +) + +# --- Reproducibility: a reproducible data stream ------------------------------ +# 回答「第 k 步看到哪些样本」。上游只播种 random/numpy/torch;RLDS/dlimp 的四处随机点 +# (TFDS 文件级 shuffle、mixure 的 sample_from_datasets、frame shuffle、增强里的 +# tf.random.uniform)都以 seed=None 构造,TF 在构建时从全局 RNG 取值,不播种则每次 +# 启动采样顺序都不同。本组把 run seed 打进 TF 全局 RNG,并显式送给 mixture 采样与 +# frame shuffle —— 也就是已验证分支显式播种的那两处。 +# +# frame shuffle 锚在 `dlimp.DLataset.shuffle`,**不是** `tf.data.Dataset.shuffle`: +# TFDS 的文件级 shuffle(`instruction_ds.shuffle(seed=read_config.shuffle_seed)`, +# dlimp 不传 shuffle_seed)走的是后者,一起改掉会改变 shard 读取顺序,曲线就与已验证 +# 分支对不上了(eager 下 seed=None 派生自全局种子,同种子即同顺序)。 +# `dlimp.DLataset.shuffle` 恰好只命中 `make_interleaved_dataset` 里那次 frame shuffle。 +# +# `prismatic.util.set_global_seed` 这个 alias 是必需的:`prismatic/util/__init__.py` +# 在导入时把 set_global_seed 重新绑定进 `prismatic.util`,而 `vla-scripts/train.py` +# 与 `scripts/pretrain.py` 都是从 `prismatic.util` 导入它的。 +# `make_interleaved_dataset` 的两个 alias 同样必需:`rlds/__init__.py` 重新导出了它, +# 而 `datasets.py` 在模块导入时把该名字绑进了自己的命名空间(`RLDSDataset.make_dataset` +# 调用的就是这个绑定)。 +# options: explicit_seeds(默认 true)、hold_shuffle_permutation(默认 true)。 +# +# 与 `openvla.dataloader.spawn` 的配合:spawn worker 里框架 patch 不生效(bootstrap 会 +# 跳过 `python -c` 辅助进程),因此每个 worker 的数据流种子由 spawn 重建路径回调 +# `reproducibility.worker_seeded_dataset_class` 安装;没有本组发布的环境标记时它是空操作。 +DATA_ORDER = group( + "openvla.reproducibility.data_order", + wrap( + target="prismatic.util.torch_utils.set_global_seed", + aliases=("prismatic.util.set_global_seed",), + replacement="turbo_physai.optimizations.models.openvla.reproducibility.set_global_seed_wrapper", + ), + wrap( + target="prismatic.util.torch_utils.worker_init_function", + replacement="turbo_physai.optimizations.models.openvla.reproducibility.worker_init_function_wrapper", + ), + wrap( + target="prismatic.vla.datasets.rlds.dataset.make_dataset_from_rlds", + replacement="turbo_physai.optimizations.models.openvla.reproducibility.make_dataset_from_rlds_wrapper", + ), + wrap( + target="prismatic.vla.datasets.rlds.dataset.make_interleaved_dataset", + aliases=( + "prismatic.vla.datasets.rlds.make_interleaved_dataset", + "prismatic.vla.datasets.datasets.make_interleaved_dataset", + ), + replacement="turbo_physai.optimizations.models.openvla.reproducibility.make_interleaved_dataset_wrapper", + ), + wrap( + target="dlimp.DLataset.sample_from_datasets", + replacement="turbo_physai.optimizations.models.openvla.reproducibility.sample_from_datasets_wrapper", + ), + wrap( + target="dlimp.DLataset.shuffle", + replacement="turbo_physai.optimizations.models.openvla.reproducibility.dlimp_shuffle_wrapper", + ), +) + +# --- Reproducibility: consistent kernel selection ----------------------------- +# 回答「同样的样本算出同样的数」。种子管不到 kernel 选择:cudnn/MIOpen 的 autotune 可能 +# 挑中不同的卷积实现、TF32 截断 fp32 尾数、部分反向归约顺序不定。本组固定这些选择。 +# +# 锚点选 `PrismaticVLM.__init__` 的理由:它是训练循环之前最后碰全局 RNG 的地方 +# (内部 `torch.manual_seed(vision_backbone.embed_dim)` 会覆盖 run seed,本组在构造后 +# 重新落实), 是每个训练/评测入口都必然经过的点, 且与其它组不冲突(set_global_seed 归 +# data_order、run_vla_training 归 gc.freeze、run_setup 归 compile.fsdp1、DataLoader 归 +# dataloader.spawn)。容器级确定性(MIOpen / rocBLAS)必须由 RuntimeConfig 在进程启动前 +# 设置,不在本组内。 +# options: cudnn_deterministic(true)、benchmark(false)、allow_tf32(false)、 +# deterministic_algorithms(true)、strict(false)。 +DETERMINISM = group( + "openvla.reproducibility.determinism", + wrap( + target="prismatic.models.vlms.prismatic.PrismaticVLM.__init__", + replacement="turbo_physai.optimizations.models.openvla.reproducibility.prismatic_vlm_init_wrapper", + ), +) + +# --- Python GC freeze before the training loop ------------------------------- +GC_FREEZE = group( + "openvla.gc.freeze", + wrap( + target=( + "prismatic.training.strategies.base_strategy." + "TrainingStrategy.run_vla_training" + ), + replacement="turbo_physai.optimizations.models.openvla.gc_freeze.gc_freeze_wrapper", + ), +) diff --git a/turbo_physai/optimizations/models/openvla/compile_fsdp1.py b/turbo_physai/optimizations/models/openvla/compile_fsdp1.py new file mode 100644 index 0000000..066848c --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/compile_fsdp1.py @@ -0,0 +1,377 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""OpenVLA FSDP1 per-layer ``torch.compile`` wrappers for TurboPhysAI. + +What this does +-------------- +Baseline ``prismatic.training.strategies.fsdp.FSDPStrategy`` has *no* +``torch.compile`` support at all. Whole-model ``torch.compile`` would trace +``LlamaModel.forward`` whose layer loop calls into ``torch.utils.checkpoint`` +(the checkpoint HOP); Dynamo cannot introspect that HOP inside a traced loop and +skips the whole frame (``convert_frame.py:854``), so the LLM silently falls back +to eager. + +This module delivers the FSDP1 per-layer compile flavour (the same idea as +``compileable_fsdp.py`` in the working tree) by **wrapping two methods that DO +exist on the baseline FSDPStrategy**: + + * ``FSDPStrategy.run_setup`` -> ``run_setup_wrapper`` (compile branch) + * ``FSDPStrategy.save_checkpoint``-> ``save_checkpoint_wrapper`` + (strips ``_orig_mod``) + +Each wrapper reads the ``openvla.compile.fsdp1`` Group ``options``. The +on branch (when ``options.compile`` is true) reproduces the baseline +``run_setup`` body but: + + 1. keeps ``buffer_dtype=None`` (avoids the Dynamo fake-tensor ``set_`` dtype + error under bf16, PyTorch #152162/#161153); + 2. instead of ``apply_activation_checkpointing`` (which puts the checkpoint HOP + *inside* the compiled unit), compiles each FSDP1-wrapped LLM transformer + layer individually and puts ``checkpoint_wrapper`` OUTSIDE the compiled unit; + 3. whole-model-compiles each FSDP1-wrapped vision tower (no checkpoint HOP + lives inside a tower, so a whole-model compile is safe there). + +Switch +------ +Whether the compile branch is used is decided by the ``openvla.compile.fsdp1`` +Group ``options`` in the OptimizationConfig / recipe -- **not** by any +environment variable: + +* ``options.compile`` (bool, default ``False``): enable per-layer + ``torch.compile``. +* ``options.mode`` (str, default ``"default"``): the ``torch.compile`` mode + (``default`` / ``reduce-overhead`` / ``max-autotune``). + +Each replacement is declared as a ``wrap`` whose factory reads those options. +When ``options.compile`` is ``False`` the factory returns the *original* +baseline method untouched, so the run is exactly the un-optimized baseline +(nothing is compiled and nothing is patched). Compile is therefore an explicit +opt-in via config rather than a silent environment default. + +Note on imports +--------------- +This module is only imported when TurboPhysAI applies the Group, i.e. inside the +real training subprocess launched by ``turbo-physai run`` (baseline on +``sys.path``), so top-level ``import prismatic`` would be safe here. In practice +the branch bodies only ever act on an existing ``FSDPStrategy`` instance via +``self`` + torch/fsdp/timm, so they do not need any ``prismatic`` symbol; we +therefore keep the module importable anywhere (no prismatic dependency at all), +which sidesteps interpreter-startup ordering entirely. +""" + +from __future__ import annotations + +import logging +import math +import os +from collections import OrderedDict +from pathlib import Path +from typing import Any, Callable, Optional + +import torch +import torch.distributed as dist +import torch.nn as nn +from timm.models.vision_transformer import VisionTransformer +from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( + CheckpointImpl, + checkpoint_wrapper, +) +from torch.distributed.fsdp import ( + FullStateDictConfig, + MixedPrecision, + StateDictType, +) +from torch.distributed.fsdp import FullyShardedDataParallel as FSDP +from transformers.optimization import get_constant_schedule, get_cosine_schedule_with_warmup + +# A plain logger (rank-agnostic; informational only). The baseline overwatch +# logger is not needed to reproduce the compile behaviour. +logger = logging.getLogger("openvla.compile.fsdp1") + + +# --- runtime switch (config-driven) --------------------------------------- +# Compile is OFF unless the Group explicitly opts in via its ``options``. +# See the module docstring for ``compile`` / ``mode`` semantics. + + +def _compile_options(options) -> tuple[bool, str]: + """Return ``(enabled, mode)`` from the Group ``options`` mapping.""" + options = dict(options or {}) + enabled = bool(options.get("compile", False)) + mode = str(options.get("mode", "default") or "default") + return enabled, mode + + +def _fsdp_class() -> type: + """Resolve the FSDP1 class to *construct* with, at call time. + + The module-level ``FSDP`` name above is bound when this module is imported, + and TurboPhysAI imports every replacement module during its preparation + phase -- i.e. **before** any Group is installed. So that binding is always + the *unpatched* class and must not be used to construct: the + ``openvla.fsdp.prefetch`` Group replaces + ``torch.distributed.fsdp.FullyShardedDataParallel`` with a subclass that + injects ``limit_all_gathers`` / ``forward_prefetch`` / ``backward_prefetch``, + and only a dynamic lookup sees it. Without this, enabling both Groups would + silently drop the prefetch optimization whenever compile is on. + + The module-level ``FSDP`` binding is still the right object for the + ``isinstance`` / ``assert isinstance`` checks below: an instance built from + the subclass satisfies them too, whether or not the patch is installed. + """ + from torch.distributed.fsdp import FullyShardedDataParallel + + return FullyShardedDataParallel + + +# --- per-layer compile helpers --------------------------------------------- +def _replace_child(module: nn.Module, name: str, child: nn.Module) -> None: + """Replace a child submodule (mirrors ``apply_activation_checkpointing``).""" + if isinstance(module, nn.Sequential): + module[int(name)] = child + else: + setattr(module, name, child) + + +def _per_layer_compile_and_checkpoint( + module: nn.Module, + layer_cls: type, + compile_mode: Optional[str], + checkpoint_module: bool, +) -> None: + """DFS: for every FSDP1-wrapped submodule of ``layer_cls`` compile + checkpoint.""" + for name, child in list(module.named_children()): + if isinstance(child, FSDP) and isinstance(child._fsdp_wrapped_module, layer_cls): + if compile_mode is not None: + child = torch.compile(child, mode=compile_mode, fullgraph=False, dynamic=True) + if checkpoint_module: + child = checkpoint_wrapper(child, checkpoint_impl=CheckpointImpl.NO_REENTRANT) + _replace_child(module, name, child) + else: + _per_layer_compile_and_checkpoint(child, layer_cls, compile_mode, checkpoint_module) + + +def _compile_whole_vision_towers(vision_backbone: nn.Module, compile_mode: Optional[str]) -> None: + """Whole-model ``torch.compile`` each FSDP-wrapped ViT tower.""" + if compile_mode is None: + return + for name, child in list(vision_backbone.named_children()): + if isinstance(child, FSDP) and isinstance(child._fsdp_wrapped_module, VisionTransformer): + _replace_child( + vision_backbone, + name, + torch.compile(child, mode=compile_mode, fullgraph=False, dynamic=True), + ) + + +# --- FSDPStrategy.run_setup wrapper ----------------------------------------- +def run_setup_wrapper(original: Callable[..., Any], options) -> Callable[..., Any]: + """wrap factory for ``FSDPStrategy.run_setup``. + + ``options.compile`` False -> return the baseline method untouched (clean + baseline, nothing compiled). True -> return a ``run_setup`` that reproduces + the baseline body and additionally applies per-layer ``torch.compile`` using + ``options.mode``. + """ + enabled, compile_mode = _compile_options(options) + if not enabled: + return original + + def run_setup(self, run_dir: Path, n_train_examples: int) -> None: + """FSDP run_setup with per-layer torch.compile (the on branch). + + Reproduces baseline ``FSDPStrategy.run_setup`` but sets + ``buffer_dtype=None`` and wraps each FSDP1 LLM layer with + ``torch.compile`` then ``checkpoint_wrapper`` OUTSIDE the compiled unit, + and whole-model-compiles the vision towers. + """ + vlm_fsdp_wrapping_policy = self.vlm.get_fsdp_wrapping_policy() + + # Mixed precision policy. + if self.enable_mixed_precision_training and self.mixed_precision_dtype == torch.bfloat16: + reduce_buffer_dtype = torch.bfloat16 if not self.reduce_in_full_precision else torch.float32 + # compile on -> buffer_dtype=None to avoid Dynamo fake-tensor set_ error. + fsdp_precision_policy = MixedPrecision( + param_dtype=torch.bfloat16, + reduce_dtype=reduce_buffer_dtype, + buffer_dtype=None, + ) + if self.stage not in {"full-finetune", "vla-full-train", "vla-sandwich-train"}: + logger.info("Casting Vision Backbone to *Half Precision* via `.to(dtype=...)`") + self.vlm.vision_backbone.to(dtype=self.vlm.vision_backbone.half_precision_dtype) + else: + fsdp_precision_policy = MixedPrecision( + param_dtype=torch.float32, reduce_dtype=torch.float32, buffer_dtype=torch.float32 + ) + + # FSDP wrap (keeps baseline sharding semantics untouched). + # `_fsdp_class()` -- resolved at call time so `openvla.fsdp.prefetch` + # (installed on the same class attribute) is honoured here as well. + # + # NOTE: `limit_all_gathers` is deliberately NOT passed here. Baseline + # semantics are preserved either way because torch's default is `True`, + # whereas passing it explicitly would make this call site the "explicit + # caller" and win over the `openvla.fsdp.prefetch` injection (the + # wrapper uses `kwargs.setdefault`), silently keeping the all-gather + # rate limiter on for the compiled path. + self.vlm = _fsdp_class()( + self.vlm, + auto_wrap_policy=vlm_fsdp_wrapping_policy, + mixed_precision=fsdp_precision_policy, + sharding_strategy=self.fsdp_sharding_strategy, + device_id=torch.cuda.current_device(), + use_orig_params=True, + ) + + # Per-layer compile + gradient checkpointing OUTSIDE the compiled unit. + # (FSDP(layer) -> torch.compile(fsdp_layer) -> checkpoint_wrapper(compiled)). + _per_layer_compile_and_checkpoint( + module=self.vlm, + layer_cls=self.llm_transformer_layer_cls, + compile_mode=compile_mode, + checkpoint_module=self.enable_gradient_checkpointing, + ) + + # Whole-model compile of the FSDP-wrapped vision towers. + _compile_whole_vision_towers( + vision_backbone=self.vlm.vision_backbone, + compile_mode=compile_mode, + ) + + dist.barrier() + + # Optimizer & LR scheduler (torch native AdamW, same as baseline FSDPStrategy). + # Resolve `torch.optim.AdamW` dynamically at call time (not a module-level + # `from torch.optim import AdamW` binding) so the `openvla.adamw.fused` wrap + # factory -- installed on `torch.optim.AdamW` at startup -- is honoured here. + n_train_examples = math.ceil(n_train_examples / self.global_batch_size) * self.global_batch_size + if self.max_steps is None: + num_training_steps = (n_train_examples * self.epochs) // self.global_batch_size + else: + num_training_steps = self.max_steps + + if self.lr_scheduler_type == "linear-warmup+cosine-decay": + num_warmup_steps = int(num_training_steps * self.warmup_ratio) + decay, no_decay = [], [] + for name, param in self.vlm.named_parameters(): + if not param.requires_grad: + continue + if param.ndim <= 1 or name.endswith(".bias"): + no_decay.append(param) + else: + decay.append(param) + groups = [{"params": decay, "weight_decay": self.weight_decay}, {"params": no_decay, "weight_decay": 0.0}] + self.optimizer = torch.optim.AdamW(groups, lr=self.learning_rate) + self.lr_scheduler = get_cosine_schedule_with_warmup(self.optimizer, num_warmup_steps, num_training_steps) + for param_group in self.optimizer.param_groups: + param_group["lr"] = 0.0 + elif self.lr_scheduler_type == "constant": + num_warmup_steps = 0 + decay, no_decay = [], [] + for name, param in self.vlm.named_parameters(): + if not param.requires_grad: + continue + if param.ndim <= 1 or name.endswith(".bias"): + no_decay.append(param) + else: + decay.append(param) + groups = [{"params": decay, "weight_decay": self.weight_decay}, {"params": no_decay, "weight_decay": 0.0}] + self.optimizer = torch.optim.AdamW(groups, lr=self.learning_rate) + self.lr_scheduler = get_constant_schedule(self.optimizer) + else: + raise ValueError(f"Learning Rate Schedule with type `{self.lr_scheduler_type}` is not supported!") + + logger.info( + "FSDP1 compile run_setup done (mode=`%s`, sharding=`%s`). VLM FSDP = %s", + compile_mode, + self.fsdp_sharding_strategy, + type(self.vlm).__name__, + ) + + return run_setup + + +# --- FSDPStrategy.save_checkpoint wrapper ----------------------------------- +def _partition_state_dicts( + full_state_dict, + module_keys, + *, + strip_orig_mod: bool = True, +): + """Split a flattened (FSDP full) state dict into per-module-key OrderedDicts. + + ``module_keys`` are the top-level name prefixes of the modules to keep + (``trainable_module_keys`` or ``all_module_keys``). When ``strip_orig_mod`` + is true, ``._orig_mod.`` segments introduced by per-layer ``torch.compile`` + are removed so the produced checkpoint keys match the plain FSDP format. + + This is exactly the mapping used by ``save_checkpoint``; exposing it makes + the compiled-vs-plain weight consistency directly testable without FSDP/CUDA. + """ + model_state_dicts = {mkey: OrderedDict() for mkey in module_keys} + for key, param in full_state_dict.items(): + if strip_orig_mod: + key = key.replace("._orig_mod.", ".") + for mkey in module_keys: + if key.startswith(prefix := f"{mkey}."): + model_state_dicts[mkey][key.removeprefix(prefix)] = param + break + return model_state_dicts + + +def save_checkpoint_wrapper(original: Callable[..., Any], options) -> Callable[..., Any]: + """wrap factory for ``FSDPStrategy.save_checkpoint``. + + ``options.compile`` False -> return the baseline method untouched (clean + baseline). True -> return a ``save_checkpoint`` that strips the + per-layer-compile ``._orig_mod.`` key segments before saving. + """ + enabled, _compile_mode = _compile_options(options) + if not enabled: + return original + + def save_checkpoint( + self, + run_dir: Path, + global_step: int, + epoch: int, + train_loss: Optional[float] = None, + only_trainable: bool = True, + ) -> None: + """Baseline save_checkpoint + strip per-layer-compile ``._orig_mod.`` key segments. + + After per-layer ``torch.compile`` each compiled layer / tower is an + ``OptimizedModule`` whose ``_orig_mod`` submodule shows up in state_dict keys + (``...layers.0._orig_mod.self_attn...``). We strip that segment so the saved + checkpoint format stays byte-identical to the plain FSDP strategy. + """ + vlm = getattr(self.vlm, "_orig_mod", self.vlm) + assert isinstance(vlm, FSDP), "save_checkpoint assumes VLM is already wrapped in FSDP!" + + with FSDP.state_dict_type(vlm, self.fsdp_state_dict_type, self.fsdp_save_policy): + full_vlm_state_dict = vlm.state_dict() + model_state_dicts = _partition_state_dicts( + full_vlm_state_dict, + self.trainable_module_keys if only_trainable else self.all_module_keys, + strip_orig_mod=True, + ) + + # Baseline uses a module-level overwatch logger; we reproduce rank0-only + # saving via the distributed backend / RANK env instead. + try: + is_rank_zero = dist.get_rank() == 0 if dist.is_initialized() else True + except Exception: + is_rank_zero = int(os.environ.get("RANK", "0")) == 0 + + if is_rank_zero: + checkpoint_dir = run_dir / "checkpoints" + if train_loss is None: + checkpoint_path = checkpoint_dir / f"step-{global_step:06d}-epoch-{epoch:02d}-loss=inf.pt" + else: + checkpoint_path = ( + checkpoint_dir / f"step-{global_step:06d}-epoch-{epoch:02d}-loss={train_loss:.4f}.pt" + ) + torch.save({"model": model_state_dicts}, checkpoint_path) + + return save_checkpoint diff --git a/turbo_physai/optimizations/models/openvla/fixbf16support.py b/turbo_physai/optimizations/models/openvla/fixbf16support.py new file mode 100644 index 0000000..e997127 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/fixbf16support.py @@ -0,0 +1,36 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any + +import torch + + +def check_bloat16_supported() -> bool: + try: + import packaging.version + import torch.cuda.nccl as nccl + import torch.distributed as dist + + if torch.version.cuda: + return ( + torch.cuda.is_bf16_supported() + and (packaging.version.parse(torch.version.cuda).release >= (11, 0)) + and dist.is_nccl_available() + and (nccl.version() >= (2, 10)) + ) + elif torch.version.hip: + return ( + torch.cuda.is_available() + and torch.cuda.is_bf16_supported() + and dist.is_nccl_available() + and (nccl.version() >= (2, 10)) + ) + else: + return False + + except Exception: + return False \ No newline at end of file diff --git a/turbo_physai/optimizations/models/openvla/fsdp_prefetch.py b/turbo_physai/optimizations/models/openvla/fsdp_prefetch.py new file mode 100644 index 0000000..8d47c2d --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/fsdp_prefetch.py @@ -0,0 +1,88 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +from __future__ import annotations + +import logging +from collections.abc import Mapping +from typing import Any, Optional + +from torch.distributed.fsdp import BackwardPrefetch + +logger = logging.getLogger("openvla.fsdp.prefetch") + +_BACKWARD_PREFETCH = { + "pre": BackwardPrefetch.BACKWARD_PRE, + "post": BackwardPrefetch.BACKWARD_POST, + "none": None, +} + + +def _prefetch_options(options: Optional[Mapping[str, Any]]) -> tuple[bool, bool, str]: + """Parse the Group ``options`` into ``(limit_all_gathers, forward_prefetch, backward_prefetch)``. + + Defaults are the overlap-optimized values: the Group being enabled *is* the + opt-in, so an empty ``options`` mapping means "apply the optimization". + """ + options = dict(options or {}) + limit_all_gathers = bool(options.get("limit_all_gathers", False)) + forward_prefetch = bool(options.get("forward_prefetch", True)) + + raw_backward_prefetch = options.get("backward_prefetch", "pre") + if raw_backward_prefetch is None: + raw_backward_prefetch = "none" + backward_prefetch = str(raw_backward_prefetch).lower() + if backward_prefetch not in _BACKWARD_PREFETCH: + raise ValueError( + "openvla.fsdp.prefetch: options.backward_prefetch must be one of " + f"{sorted(_BACKWARD_PREFETCH)}, got {raw_backward_prefetch!r}" + ) + return limit_all_gathers, forward_prefetch, backward_prefetch + + +def fsdp_prefetch_wrapper(original: Any, options: Optional[Mapping[str, Any]] = None): + """Wrapper factory ``(original, options) -> FSDP class with overlapped collectives``. + + ``original`` is ``torch.distributed.fsdp.FullyShardedDataParallel``. The + returned class injects the three communication keyword defaults into every + construction while leaving the submitted values untouched. + """ + if not isinstance(original, type): + return original + + limit_all_gathers, forward_prefetch, backward_prefetch = _prefetch_options(options) + backward_prefetch_value = _BACKWARD_PREFETCH[backward_prefetch] + + def prefetch_init(self: Any, *args: Any, **kwargs: Any) -> None: + kwargs.setdefault("limit_all_gathers", limit_all_gathers) + kwargs.setdefault("forward_prefetch", forward_prefetch) + kwargs.setdefault("backward_prefetch", backward_prefetch_value) + original.__init__(self, *args, **kwargs) + + # Build the subclass through `type()` so the installed class keeps the exact + # original name (torch/dynamo internals and log lines that key off + # `type(module).__name__` stay indistinguishable from the unpatched run). + replacement = type( + original.__name__, + (original,), + { + "__init__": prefetch_init, + "__doc__": ( + f"{original.__doc__ or ''}\n\n" + "[openvla.fsdp.prefetch] Injected FSDP1 communication defaults: " + f"limit_all_gathers={limit_all_gathers}, " + f"forward_prefetch={forward_prefetch}, " + f"backward_prefetch={backward_prefetch!r}." + ), + "__module__": original.__module__, + }, + ) + + logger.info( + "FSDP1 wrapped with `openvla.fsdp.prefetch` (limit_all_gathers=%s, " + "forward_prefetch=%s, backward_prefetch=%s)", + limit_all_gathers, + forward_prefetch, + backward_prefetch, + ) + return replacement diff --git a/turbo_physai/optimizations/models/openvla/fusedAdamW.py b/turbo_physai/optimizations/models/openvla/fusedAdamW.py new file mode 100644 index 0000000..57860fc --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/fusedAdamW.py @@ -0,0 +1,15 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any + + +def adamw_fused_wrapper(original: Any, options: Mapping[str, Any]): + def adamw_factory(*args: Any, **kwargs: Any): + kwargs.setdefault("fused", True) + return original(*args, **kwargs) + + return adamw_factory \ No newline at end of file diff --git a/turbo_physai/optimizations/models/openvla/gc_freeze.py b/turbo_physai/optimizations/models/openvla/gc_freeze.py new file mode 100644 index 0000000..ee6a918 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/gc_freeze.py @@ -0,0 +1,51 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +from __future__ import annotations + +import functools +import gc +from collections.abc import Mapping +from typing import Any, Callable, Optional + + +def _log(message: str) -> None: + """Best-effort rank-zero log; never affects the optimization itself.""" + + try: + import overwatch + + if overwatch.is_rank_zero(): + overwatch.info(message) + except Exception: # noqa: BLE001 - logging must never break training + pass + + +def freeze_gc() -> int: + """Move everything currently tracked into the permanent generation. + + Returns the number of objects now parked there. Safe to call more than + once: ``gc.freeze()`` is additive, and a later call additionally parks + whatever was created in the meantime. + """ + + gc.freeze() + return gc.get_freeze_count() + + +def gc_freeze_wrapper( + original: Callable, options: Optional[Mapping[str, Any]] = None +) -> Callable: + + del options + + @functools.wraps(original) + def run_vla_training(self, *args, **kwargs): + frozen = freeze_gc() + _log( + "Python GC policy ENABLED =>> freeze=1; " + f"{frozen} objects moved to the permanent generation before the training loop" + ) + return original(self, *args, **kwargs) + + return run_vla_training diff --git a/turbo_physai/optimizations/models/openvla/reproducibility.py b/turbo_physai/optimizations/models/openvla/reproducibility.py new file mode 100644 index 0000000..b60b95d --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/reproducibility.py @@ -0,0 +1,673 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""OpenVLA reproducibility: a reproducible data stream and consistent kernels. + +This module implements two independent Groups. They answer two different +questions and they complement each other: + +``openvla.reproducibility.data_order`` -- *which* samples does a step see? + Upstream OpenVLA seeds ``random`` / ``numpy`` / ``torch`` in + ``prismatic.util.torch_utils.set_global_seed`` and stops there, but the RLDS + pipeline (``dlimp`` / ``tfds``) constructs several ops with ``seed=None``: + ``Dataset.shuffle()`` and ``tf.data.Dataset.sample_from_datasets()`` inside + ``make_interleaved_dataset``, TFDS' file-level ``shuffle_files=True`` (dlimp + passes no ``shuffle_seed``) and the ``tf.random.uniform`` draw in the frame + transform. TensorFlow resolves ``seed=None`` from its *global* RNG when the op + is constructed, so without seeding it the sample order is re-randomized on + every launch and two runs cannot be compared step by step. + + The Group therefore pins the run seed before the RLDS graph is built, and + threads it explicitly into the mixture sampling and the *frame* shuffle -- the + two places the validated fork seeds explicitly. TFDS' file-level shuffle is + deliberately left to the global RNG, exactly as the fork leaves it; see + :func:`dlimp_shuffle_wrapper` for why pinning that one too would *change* the + batch stream instead of reproducing the fork's. + +``openvla.reproducibility.determinism`` -- do the same samples *compute* the same? + Seeds say nothing about kernel selection: (c)uDNN/MIOpen may autotune to a + different, non-deterministic convolution, TF32 truncates the fp32 mantissa, + and some backward reductions accumulate in a non-deterministic order. This + Group pins those choices, and re-asserts the run seed after ``PrismaticVLM`` + construction -- whose ``__init__`` calls + ``torch.manual_seed(vision_backbone.embed_dim)`` and would otherwise *become* + the effective torch seed of the run. + + It anchors on ``PrismaticVLM.__init__`` because that constructor is the last + thing to touch the global RNG before the training loop, the first point every + training/eval entry point has certainly reached, and the earliest point at + which the run seed is known (upstream sets it before loading the model). It + also collides with no other Group: ``set_global_seed`` belongs to + ``data_order``, ``run_vla_training`` to ``openvla.gc.freeze``, ``run_setup`` + to ``openvla.compile.fsdp1`` and the ``DataLoader`` to + ``openvla.dataloader.spawn``. + +Only both Groups together give a bit-reproducible loss. ``data_order`` alone +fixes the samples but leaves kernel noise (the loss then agrees to a few +significant digits); ``determinism`` alone fixes the arithmetic while the samples +still differ from run to run. Container-level determinism +(``MIOPEN_DEBUG_CONVOLUTION_DETERMINISTIC``, ``MIOPEN_FIND_MODE``, the rocBLAS +tuning library) is deliberately *not* here: those variables are read when the +backend initializes and have to be set by the RuntimeConfig before the +interpreter starts. + +Where the run seed travels +-------------------------- +``set_global_seed`` publishes the seed in-process and in the environment +(``TURBO_PHYSAI_OPENVLA_SEED``). The environment copy is what makes spawn +DataLoader workers work: they re-import the model stack, so the framework's +replacements are not installed there, and the marker is inherited instead. When +only ``data_order`` is enabled the seed is still found, through the upstream +``EXPERIMENT_GLOBAL_SEED`` marker that ``prismatic`` itself writes. With no seed +published at all every wrapper is a pass-through, i.e. the baseline behaviour is +bit-identical. + +Options +------- +``openvla.reproducibility.data_order`` + ``explicit_seeds`` (bool, default ``true``) + Thread the run seed into ``dlimp``'s mixture sampling and into the + ``tf.data`` frame shuffle, on top of the global TF seed. + ``hold_shuffle_permutation`` (bool, default ``true``) + Pass ``reshuffle_each_iteration=False`` on the seeded frame shuffle, so + the permutation is a function of the seed alone. + +``openvla.reproducibility.determinism`` + The five options mirror the torch attributes they set; their defaults are the + bundle the validated fork applies behind ``OPENVLA_DETERMINISTIC=1``: + ``cudnn_deterministic`` (true), ``benchmark`` (false), ``allow_tf32`` (false), + ``deterministic_algorithms`` (true) and ``strict`` (false). ``strict`` makes + ``use_deterministic_algorithms`` hard-fail instead of warn on the few ops + that have no deterministic kernel on ROCm. + +Known limits +------------ +* ``--image_aug`` is not reproducible: the augmentation seed is drawn inside a + ``tf.data`` map running with 16 parallel calls, so the values a frame receives + depend on the order in which the replicas consume the stateful op. +* Bit-for-bit equality is a same-machine, same-world-size property. Different + device counts shard and reduce differently, and the RLDS normalization + statistics can be recomputed with a different file read order. +* ``metrics.jsonl`` is written per rank and never reduced across ranks, so a + comparison is between the same rank of two runs. + +Module import is side-effect free and does not import torch / tensorflow / +prismatic: every dependency is imported inside the function that needs it. +""" + +from __future__ import annotations + +import functools +import os +from collections.abc import Mapping +from typing import Any, Callable, Optional + +__all__ = [ + "RUN_SEED_ENV", + "active_run_seed", + "apply_determinism", + "dlimp_shuffle_wrapper", + "make_dataset_from_rlds_wrapper", + "make_interleaved_dataset_wrapper", + "prismatic_vlm_init_wrapper", + "reseed_host_rngs", + "run_seed", + "sample_from_datasets_wrapper", + "seed_tensorflow", + "set_global_seed_wrapper", + "set_run_seed", + "worker_init_function_wrapper", + "worker_seeded_dataset_class", +] + + +#: Seed marker of these Groups; unlike a process global it survives ``spawn``. +RUN_SEED_ENV = "TURBO_PHYSAI_OPENVLA_SEED" +#: Seed marker upstream ``prismatic`` writes in its own ``set_global_seed``. +UPSTREAM_SEED_ENV = "EXPERIMENT_GLOBAL_SEED" + +_RUN_SEED: Optional[int] = None + +#: The wrapped ``worker_init_function`` built by this module's factory. +#: ``set_global_seed`` hands it to the DataLoader, so the two members of the +#: ``data_order`` Group have to agree on one object. +_WORKER_INIT: Optional[Callable[[int], None]] = None + + +# --- options and the run seed ------------------------------------------------ + +def _option(options: Optional[Mapping[str, Any]], name: str, default: Any) -> Any: + """Read one Group option, with a default for a missing or empty mapping.""" + + return dict(options or {}).get(name, default) + + +def _option_enabled(value: Any, default: bool) -> bool: + """Parse a boolean Group option (YAML bool, or 1/0, true/false, yes/no, on/off).""" + + if value is None: + return default + if isinstance(value, bool): + return value + if isinstance(value, (int, float)): + return bool(value) + text = str(value).strip().lower() + if text in {"1", "true", "yes", "on"}: + return True + if text in {"0", "false", "no", "off"}: + return False + return default + + +def _env_seed(name: str) -> Optional[int]: + """Read an integer seed from the environment; ``None`` when absent or invalid.""" + + raw = os.environ.get(name) + if raw is None or not raw.strip(): + return None + try: + return int(raw) + except ValueError: + return None + + +def set_run_seed(seed: int) -> int: + """Record the run seed in-process and publish it to spawned children. + + The environment copy is not a convenience: a spawn DataLoader worker + re-imports the model stack and therefore never sees this module's globals. + """ + + global _RUN_SEED + _RUN_SEED = int(seed) + os.environ[RUN_SEED_ENV] = str(_RUN_SEED) + return _RUN_SEED + + +def run_seed() -> Optional[int]: + """Effective run seed, or ``None`` when nothing has been seeded yet. + + Resolution order: the in-process record (written by this Group's + ``set_global_seed`` wrapper), then this Group's environment marker, then the + upstream ``EXPERIMENT_GLOBAL_SEED`` marker that ``prismatic`` itself writes. + The last step is what lets ``data_order`` work with the ``seed`` half + disabled and what lets ``determinism`` re-seed a run that only seeded the + host libraries. + """ + + if _RUN_SEED is not None: + return _RUN_SEED + for name in (RUN_SEED_ENV, UPSTREAM_SEED_ENV): + seed = _env_seed(name) + if seed is not None: + return seed + return None + + +def active_run_seed() -> Optional[int]: + """The seed *these* Groups published, i.e. the environment marker only. + + Used where the answer has to mean "the framework published a seed in the + parent process" -- notably :func:`worker_seeded_dataset_class`, which runs + inside a spawn worker and must stay inert when only + ``openvla.dataloader.spawn`` is enabled. + """ + + return _env_seed(RUN_SEED_ENV) + + +# --- library-level seeding --------------------------------------------------- + +def _tensorflow(): + """Return the TensorFlow module, or ``None`` when it is not installed.""" + + try: + import tensorflow + except ImportError: + return None + return tensorflow + + +def seed_tensorflow(seed: int) -> bool: + """Pin TensorFlow's *global* RNG; returns whether TF was available. + + Has to run before any TF op is constructed for the seed to matter, because TF + resolves ``seed=None`` from this global state at graph-construction time. + """ + + tensorflow = _tensorflow() + if tensorflow is None: + return False + tensorflow.random.set_seed(int(seed)) + return True + + +def apply_determinism( + cudnn_deterministic: bool = True, + benchmark: bool = False, + allow_tf32: bool = False, + deterministic_algorithms: bool = True, + strict: bool = False, +) -> bool: + """Pin kernel selection on top of the seeds; returns whether torch was available. + + Every argument mirrors the torch attribute it sets: + + * ``cudnn_deterministic`` / ``benchmark``: deterministic kernels only, and no + runtime autotuning (on ROCm this is the MIOpen path). Autotuning is the + usual reason two identical runs pick different algorithms. + * ``allow_tf32`` for matmul and (c)udnn: TF32 truncates the fp32 mantissa, so + it has to be off for a run that is compared against another vendor. + * ``deterministic_algorithms`` with ``strict=False``: deterministic kernels + and reductions, warning instead of aborting on the handful of ops ROCm has + no deterministic implementation for; ``strict=True`` hard-fails instead. + When disabled, the process setting is left untouched. + """ + + try: + import torch + except ImportError: + return False + + torch.backends.cudnn.deterministic = bool(cudnn_deterministic) + torch.backends.cudnn.benchmark = bool(benchmark) + try: + torch.backends.cuda.matmul.allow_tf32 = bool(allow_tf32) + except (AttributeError, RuntimeError): # backend without a CUDA/TF32 config + pass + try: + torch.backends.cudnn.allow_tf32 = bool(allow_tf32) + except (AttributeError, RuntimeError): # backend without a (c)udnn TF32 config + pass + if deterministic_algorithms: + torch.use_deterministic_algorithms(True, warn_only=not strict) + return True + + +def reseed_host_rngs(seed: Optional[int]) -> bool: + """Make ``seed`` the effective torch seed again; returns whether it ran. + + ``torch.manual_seed`` already forwards to ``torch.cuda.manual_seed_all``, but + the flash-attention dropout path and any per-device ``torch.Generator`` read + the per-device generators directly, so the call is made explicit. + """ + + if seed is None: + return False + try: + import torch + except ImportError: + return False + torch.manual_seed(int(seed)) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(int(seed)) + return True + + +# --- Group ``openvla.reproducibility.data_order`` ----------------------------- + +def set_global_seed_wrapper( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """``wrap`` factory for ``prismatic.util.torch_utils.set_global_seed``. + + Keeps the baseline seeding of ``random`` / ``numpy`` / ``torch`` exactly as it + is and adds the two data-stream pieces on top: + + * the run seed is recorded and published, which is what the RLDS pipeline + wrappers below and the spawn DataLoader workers read; + * TensorFlow's global RNG is seeded, which is what pins the ``seed=None`` ops + of the RLDS pipeline (see the module docstring). + + The returned ``worker_init_fn`` is this module's wrapped version, so workers + get a TensorFlow child sequence as well. + """ + + del options + + @functools.wraps(original) + def set_global_seed(seed: int, get_worker_init_fn: bool = False): + original(seed, get_worker_init_fn) # random / numpy / torch + EXPERIMENT_GLOBAL_SEED + set_run_seed(seed) + seed_tensorflow(seed) + return _WORKER_INIT if get_worker_init_fn else None + + return set_global_seed + + +def _worker_tensorflow_seed(worker_id: int) -> Optional[int]: + """TensorFlow child seed for one DataLoader worker. + + Mirrors the derivation the baseline ``worker_init_function`` already uses and + extends it with a third child sequence for TensorFlow:: + + seed_seq = SeedSequence([base_seed, worker_id, LOCAL_RANK]) + tf_seed = seed_seq.spawn(3)[2] + + ``SeedSequence.spawn(n)`` derives child ``i`` from ``(spawn_key + (i,))`` + alone, so asking for three children instead of two leaves the torch and + ``random`` seeds bit-identical to the baseline. + + Has to be called *before* the baseline ``worker_init_function``: that one + reseeds torch, after which ``torch.initial_seed()`` no longer holds the + per-worker seed the DataLoader installed. + """ + + try: + import numpy as np + import torch + except ImportError: + return None + global_rank = int(os.environ.get("LOCAL_RANK", 0)) + base_seed = int(torch.initial_seed()) - int(worker_id) + seed_sequence = np.random.SeedSequence([base_seed, int(worker_id), global_rank]) + return int(seed_sequence.spawn(3)[2].generate_state(1, dtype=np.uint64)[0]) + + +def worker_init_function_wrapper( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """``wrap`` factory for ``prismatic.util.torch_utils.worker_init_function``. + + Adds a TensorFlow child seed on top of the baseline worker seeding. This + runs *after* a spawned worker has unpickled its dataset, so it cannot + influence a TF graph that was already built: the per-worker *data stream* + seed is applied by :func:`worker_seeded_dataset_class` instead. + """ + + del options + + @functools.wraps(original) + def worker_init_function(worker_id: int) -> None: + tf_seed = _worker_tensorflow_seed(worker_id) + original(worker_id) + if tf_seed is not None: + seed_tensorflow(tf_seed) + + global _WORKER_INIT + _WORKER_INIT = worker_init_function + return worker_init_function + + +def _seed_before_pipeline_build( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """Shared factory for the two RLDS pipeline builders. + + TensorFlow captures ``seed=None`` for every dataset op (TFDS' file shuffle, + ``sample_from_datasets``, the frame ``shuffle`` and the augmentation draws) at + graph-construction time, so the run seed has to be in the global RNG *before* + the wrapped call builds those ops -- in every process that builds the graph, + including spawn DataLoader workers. + + A ``seed`` keyword supplied by the caller (the validated fork threads one) + wins over the run seed and is consumed here instead of being forwarded, + because upstream's signature has no such parameter. + """ + + del options + + @functools.wraps(original) + def builder(*args: Any, **kwargs: Any) -> Any: + seed = kwargs.pop("seed", None) + effective_seed = seed if seed is not None else run_seed() + if effective_seed is not None: + seed_tensorflow(effective_seed) + return original(*args, **kwargs) + + return builder + + +def make_dataset_from_rlds_wrapper( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """``wrap`` factory for ``make_dataset_from_rlds``. + + ``dlimp.DLataset.from_rlds(..., shuffle=True)`` shuffles the *files* through + TFDS, which resolves its seed from TensorFlow's global RNG, and the rest of + that dataset's pipeline is constructed below this call -- so this is the + earliest point of the RLDS build and the right place to pin the seed. + """ + + return _seed_before_pipeline_build(original, options) + + +def make_interleaved_dataset_wrapper( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """``wrap`` factory for ``make_interleaved_dataset``. + + The same backstop as :func:`make_dataset_from_rlds_wrapper`, applied to the + whole mixture build: per-dataset file order, the frame-level mixture sampling + and the frame shuffle all come out of the TF ops constructed below this call. + """ + + return _seed_before_pipeline_build(original, options) + + +def _accepts_keyword(function: Callable[..., Any], name: str) -> bool: + """Whether ``function`` accepts a keyword argument ``name``.""" + + try: + import inspect + + parameters = inspect.signature(function).parameters + except (TypeError, ValueError): + return False + return name in parameters or any( + parameter.kind is inspect.Parameter.VAR_KEYWORD + for parameter in parameters.values() + ) + + +def sample_from_datasets_wrapper( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """``wrap`` factory for ``dlimp.DLataset.sample_from_datasets``. + + ``make_interleaved_dataset`` calls this without a seed and dlimp forwards + ``seed=None`` to ``tf.data.Dataset.sample_from_datasets`` -- a fresh choice + stream per launch, which decides *which* dataset every frame is drawn from. + With ``explicit_seeds`` (default on) a missing seed becomes the run seed; a + seed the caller passed is never touched, and a dlimp without a ``seed`` + parameter is still called as before instead of raising ``TypeError``. + """ + + explicit_seeds = _option_enabled(_option(options, "explicit_seeds", True), True) + supports_seed = _accepts_keyword(original, "seed") + + @functools.wraps(original) + def sample_from_datasets( + datasets: Any, weights: Any = None, seed: Optional[int] = None, **kwargs: Any + ) -> Any: + if explicit_seeds and seed is None: + seed = run_seed() + if supports_seed: + kwargs["seed"] = seed + return original(datasets, weights, **kwargs) + + return sample_from_datasets + + +def dlimp_shuffle_wrapper( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """``wrap`` factory for ``dlimp.DLataset.shuffle`` -- the mixture's *frame* shuffle. + + The frame shuffle is written ``dataset.shuffle(shuffle_buffer_size)`` inside + ``make_interleaved_dataset``. With ``explicit_seeds`` (default on) a missing seed + becomes the run seed, and with ``hold_shuffle_permutation`` (default on) the + permutation is fixed to that seed instead of being redrawn per iteration. Only + ``seed=None`` calls are touched: an explicit seed from any other caller is + preserved, and with no run seed published the wrapper is a pass-through. + + Why the anchor is ``dlimp.DLataset`` and not ``tf.data.Dataset.shuffle`` + ---------------------------------------------------------------------- + ``tf.data.Dataset.shuffle`` is also the method TFDS uses for the *file-level* + shuffle of ``from_rlds(..., shuffle=True)`` + (``instruction_ds.shuffle(len(files), seed=read_config.shuffle_seed)``, with + dlimp passing no ``shuffle_seed``). Patching that method rewrites the shard + order as well, and the validated fork deliberately leaves *that* one to the + global RNG: under eager execution ``seed=None`` resolves to + ``(global_seed, random.Random(global_seed).randint(0, 2**31 - 1))``, i.e. it + depends only on the global seed and on how many seedless random ops were built + since it was set -- never on the op count of the process. Both spellings are + reproducible, but they read the shards in *different* orders, so pinning the + file shuffle here would produce a batch stream that no longer matches the + fork's (first batch included). Anchoring on ``DLataset`` hits only the frame + shuffle that ``make_interleaved_dataset`` performs, which is the one the fork + seeds explicitly. ``DLataset`` reaches this wrapper like any other + ``tf.data.Dataset`` method: its ``__getattribute__`` re-wraps the result into a + ``DLataset``. + """ + + explicit_seeds = _option_enabled(_option(options, "explicit_seeds", True), True) + hold_permutation = _option_enabled(_option(options, "hold_shuffle_permutation", True), True) + + @functools.wraps(original) + def shuffle( + self: Any, + buffer_size: Any, + seed: Optional[int] = None, + reshuffle_each_iteration: Optional[bool] = None, + **kwargs: Any, + ) -> Any: + if explicit_seeds and seed is None: + effective_seed = run_seed() + if effective_seed is not None: + seed = int(effective_seed) + if hold_permutation and reshuffle_each_iteration is None: + reshuffle_each_iteration = False + return original( + self, + buffer_size, + seed=seed, + reshuffle_each_iteration=reshuffle_each_iteration, + **kwargs, + ) + + return shuffle + + +# --- Group ``openvla.reproducibility.determinism`` ---------------------------- + +def prismatic_vlm_init_wrapper( + original: Callable[..., Any], options: Optional[Mapping[str, Any]] = None +) -> Callable[..., Any]: + """``wrap`` factory for ``prismatic.models.vlms.prismatic.PrismaticVLM.__init__``. + + ``PrismaticVLM.__init__`` calls ``torch.manual_seed(vision_backbone.embed_dim)`` + so that the projector initialization is deterministic -- which *overwrites* + the run seed, and would otherwise make the ViT embedding dim the effective + torch seed of the run. Re-asserting the run seed here is what makes + ``--seed`` mean what it says, and doing it in the constructor (rather than + after model loading in the training script) also covers evaluation scripts + and any other entry point that builds a VLM. + + Nothing between the construction and the re-seed consumes randomness, and + nothing before the construction runs a real kernel, so neither the + determinism bundle nor the re-seed moves a draw. + """ + + settings = { + "cudnn_deterministic": _option_enabled(_option(options, "cudnn_deterministic", True), True), + "benchmark": _option_enabled(_option(options, "benchmark", False), False), + "allow_tf32": _option_enabled(_option(options, "allow_tf32", False), False), + "deterministic_algorithms": _option_enabled( + _option(options, "deterministic_algorithms", True), True + ), + "strict": _option_enabled(_option(options, "strict", False), False), + } + + @functools.wraps(original) + def __init__(self: Any, *args: Any, **kwargs: Any) -> None: + apply_determinism(**settings) + original(self, *args, **kwargs) + reseed_host_rngs(run_seed()) + + return __init__ + + +# --- Per-worker data streams (used by ``openvla.dataloader.spawn``) ----------- + +def _worker_id() -> int: + """DataLoader worker id, or ``0`` outside a worker (``spawn`` unpickling).""" + + try: + import torch.utils.data + except ImportError: + return 0 + info = torch.utils.data.get_worker_info() + return 0 if info is None else int(info.id) + + +def _install_identity(subclass: type, original: type) -> type: + """Give ``subclass`` the original class' name/doc, like the spawn wrapper does.""" + + subclass.__name__ = original.__name__ + subclass.__qualname__ = original.__qualname__ + subclass.__doc__ = original.__doc__ + return subclass + + +def worker_seeded_dataset_class(original: type, seed: Optional[int] = None) -> type: + """Return an RLDS dataset class whose TF graph is built per DataLoader worker. + + Called by the spawn rebuild path (``spawn_dataloader._reconstruct_rlds_dataset``) + inside a worker: the worker re-runs the dataset constructor there, which + happens *before* ``worker_init_fn`` and therefore too early for anything + worker-id dependent, and (without this hook) against an unseeded global RNG -- + i.e. a different data stream per worker and per launch. + + The returned subclass applies the validated fork's rule, ``run_seed + worker_id``: + + * its constructor seeds TensorFlow before ``super().__init__`` builds the + graph, and remembers the seed the live graph was built with; + * ``__iter__`` is where the worker id becomes known -- the DataLoader fetcher + calls ``iter(dataset)`` after ``worker_init_fn`` -- so the graph is rebuilt + there whenever the remembered seed is not this worker's; + * two co-existing workers each iterate a private copy of the same infinite + stream, so the ``+ worker_id`` offset is what keeps the streams disjoint: + with one shared seed every batch would be emitted twice. + + ``original`` is returned unchanged when no seed was published by these Groups + (so ``openvla.dataloader.spawn`` keeps its exact behaviour on its own) and when + the class has already been wrapped. + """ + + base_seed = active_run_seed() if seed is None else int(seed) + if base_seed is None or getattr(original, "__turbo_physai_worker_seeded__", False): + return original + + class _WorkerSeededDataset(original): + """RLDS dataset whose TensorFlow graph is rebuilt per DataLoader worker.""" + + __turbo_physai_rlds_dataset__ = True + + def __init__(self, *args: Any, **kwargs: Any) -> None: + self.__ctor_args = args + self.__ctor_kwargs = kwargs + self._turbo_physai_seed_graph() + super().__init__(*args, **kwargs) + + def _turbo_physai_seed_graph(self) -> int: + """Seed TensorFlow for this worker and remember the live graph's seed.""" + + wanted = base_seed + _worker_id() + seed_tensorflow(wanted) + self.__built_seed = wanted + return wanted + + def __iter__(self): + if self.__built_seed != base_seed + _worker_id(): + # First `iter()` inside a worker: rebuild the graph for this + # worker's seed (worker 0 reuses the graph its constructor built). + self._turbo_physai_seed_graph() + original.__init__(self, *self.__ctor_args, **self.__ctor_kwargs) + # Delegate to the base implementation instead of iterating here, so + # subclasses that yield differently (`EpisodicRLDSDataset`) keep their + # exact semantics. + return super().__iter__() + + _WorkerSeededDataset.__turbo_physai_worker_seeded__ = True + return _install_identity(_WorkerSeededDataset, original) diff --git a/turbo_physai/optimizations/models/openvla/skip_fa2_unpad.py b/turbo_physai/optimizations/models/openvla/skip_fa2_unpad.py new file mode 100644 index 0000000..f105f86 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/skip_fa2_unpad.py @@ -0,0 +1,59 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +""" +OpenVLA skip-FA2-unpad replacement for TurboPhysAI. +""" + +from __future__ import annotations + +import functools +from collections.abc import Mapping +from typing import Any, Callable, Optional + + +def _is_right_padded(attention_mask: Any) -> bool: + """True when the batch carries no left padding. + + Dropping the mask is only sound for right-padded batches. With left + padding the real tokens sit *after* the pad slots, so plain causal attention + would let them attend to those pad key/values instead of having them + unpadded away by the varlen path -- silently, and with no error. A + right-padded (or unpadded) batch has no zero in its first column. + + The ``.item()`` is a scalar D2H sync, far cheaper than the one this Group + removes: the varlen path needs a data-dependent ``torch.nonzero`` plus + ``seqlens_in_batch.max().item()`` on every step. + """ + + return attention_mask.ndim == 2 and bool(attention_mask[:, 0].all().item()) + + +def make_fast_fa2_causal_mask_wrapper( + original: Callable, options: Optional[Mapping[str, Any]] = None +) -> Callable: + """Wrapper factory ``(original, options) -> wrapped _update_causal_mask``. + + Enabling the Group installs the returned callable in place of + ``LlamaModel._update_causal_mask``. ``options`` (Group options) is accepted + for framework ``wrap`` compatibility and intentionally unused. + """ + + del options + + @functools.wraps(original) + def wrapper(self, attention_mask, input_tensor, cache_position, past_seen_tokens): + # Prefill (no KV cache) + FA2 + an explicit mask: skip the causal-mask + # materialisation entirely so FA2 never enters the varlen/unpad path. + # Right padding is a precondition, not an assumption: the mask is + # dropped, so a left-padded batch must fall back to the original. + if ( + self.config._attn_implementation == "flash_attention_2" + and attention_mask is not None + and past_seen_tokens == 0 + and _is_right_padded(attention_mask) + ): + return None + return original(self, attention_mask, input_tensor, cache_position, past_seen_tokens) + + return wrapper diff --git a/turbo_physai/optimizations/models/openvla/spawn_dataloader.py b/turbo_physai/optimizations/models/openvla/spawn_dataloader.py new file mode 100644 index 0000000..517895d --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/spawn_dataloader.py @@ -0,0 +1,270 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""OpenVLA spawn-worker DataLoader replacement for TurboPhysAI. + +Why this exists +--------------- +OpenVLA's VLA training loop (``TrainingStrategy.run_vla_training``) feeds an +RLDS (TFDS-backed) dataset through a ``DataLoader``. The baseline creates that +loader with ``num_workers=0``: the (CPU-heavy, Python) batch transform +``RLDSBatchTransform`` — image decode/resize + tokenization + label +construction — runs inline in the training process. + +The validated HCU optimization moves the RLDS data pipeline into **spawned** +DataLoader worker processes instead: + +* ``num_workers=N`` workers run the batch transform off the main process; +* workers must use the ``spawn`` multiprocessing context: RLDS keeps a + TensorFlow graph/threadpool alive in the parent, and ``fork``-ing a process + with an active TF threadpool deadlocks. ``spawn`` starts a clean process in + which the TF graph is rebuilt from scratch; +* that requires the dataset instance to be picklable, so the wrapped dataset + classes serialize their *constructor arguments* instead of the TF graph and + rebuild the graph inside the worker. + +Scope of this Group +------------------- +Only ``DataLoader`` constructions whose dataset is an RLDS dataset are touched: +they get ``num_workers=N`` (Group ``options`` key ``num_workers``, default ``1``), +the ``spawn`` context, and — when ``N > 0`` — ``pin_memory=True`` (Group +``options`` key ``pin_memory``, default ``True``). Every other DataLoader in the +process (eval loaders, non-RLDS training, HF internals, ...) keeps the exact +upstream behaviour, and so does any RLDS loader when ``num_workers`` is ``0``. + +``pin_memory=True`` is the upstream-HCU item carried by this Group: the pinned +staging buffer is filled by the worker while the main process computes, so the +main process' H2D copy reads from page-locked host memory instead of going +through a pageable copy. With ``num_workers=0`` there is no worker to overlap +with, so the option is left untouched there (baseline default ``False``). + +Three ``wrap`` members compose the Group: + +1. ``RLDSDataset`` -> spawn-pickling subclass (constructor-arg serialisation) +2. ``EpisodicRLDSDataset`` -> ditto +3. ``torch.utils.data.DataLoader`` -> subclass that forces ``num_workers=N`` + + ``spawn`` + ``pin_memory=True`` only when the dataset is an RLDS dataset (also + patches the early ``from torch.utils.data import DataLoader`` binding in + ``prismatic.training.strategies.base_strategy`` via ``aliases``). + +Module import is side-effect free and does not import torch/transformers/prismatic. + +Interaction with ``openvla.reproducibility.*`` +--------------------------------------------- +Workers rebuild the dataset by re-running its constructor, which happens *before* +``worker_init_fn`` -- too early for anything worker-id dependent. So the +per-worker data stream seed is installed from the rebuild path itself: the +reconstructor asks ``reproducibility.worker_seeded_dataset_class`` for the class +to instantiate, and that helper returns the original class unchanged unless one +of the reproducibility Groups published a run seed in the parent process (the +marker travels to the worker through the environment). This Group therefore +behaves exactly as before when those Groups are disabled. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, Callable, Optional + + +# --- helpers ----------------------------------------------------------------- + +def _rlds_num_workers(options: Optional[Mapping[str, Any]]) -> int: + """Read the ``num_workers`` Group option (default 1, >= 0).""" + options = dict(options or {}) + try: + workers = int(options.get("num_workers", 1)) + except (TypeError, ValueError): + workers = 1 + return max(0, workers) + + +def _option_enabled(value: Any, default: bool) -> bool: + """Parse a boolean Group option (YAML bool, or 1/0, true/false, yes/no, on/off).""" + if value is None: + return default + if isinstance(value, bool): + return value + if isinstance(value, (int, float)): + return bool(value) + text = str(value).strip().lower() + if text in {"1", "true", "yes", "on"}: + return True + if text in {"0", "false", "no", "off"}: + return False + return default + + +def _rlds_pin_memory(options: Optional[Mapping[str, Any]]) -> bool: + """Read the ``pin_memory`` Group option (default ``True``).""" + options = dict(options or {}) + return _option_enabled(options.get("pin_memory"), True) + + +def _looks_like_rlds_dataset(dataset: Any) -> bool: + """True when ``dataset`` is (a subclass instance of) an RLDS dataset. + + The wrapped dataset classes carry a class marker; the module/name check is a + fallback for unwrapped original ``RLDSDataset`` / ``EpisodicRLDSDataset`` + instances. No prismatic import is triggered here. + """ + cls = type(dataset) + if getattr(cls, "__turbo_physai_rlds_dataset__", False): + return True + return ( + cls.__module__ == "prismatic.vla.datasets.datasets" + and cls.__name__ in ("RLDSDataset", "EpisodicRLDSDataset") + ) + + +def _install_identity(subclass: type, original: type) -> type: + subclass.__name__ = original.__name__ + subclass.__qualname__ = original.__qualname__ + subclass.__doc__ = original.__doc__ + return subclass + + +# --- dataset spawn-pickling wrappers ----------------------------------------- + +def _spawn_picklable_dataset(original: type) -> type: + """Subclass ``original`` so instances can cross a spawn process boundary. + + An RLDS dataset owns a TensorFlow graph (``self.dataset``) that cannot be + pickled. The subclass remembers the constructor arguments and serialises + them instead (via ``__reduce__``); unpickling inside a spawn worker runs the + real constructor, which rebuilds the TF graph from scratch there. + + The base class is serialised as a **module/qualname path**, not as the class + object: pickle pickles classes by reference and verifies that the module + attribute still points at the object being pickled. Applying this Group + replaces ``prismatic.vla.datasets.datasets.RLDSDataset`` with the wrapped + subclass, so handing pickle the original class object fails with + ``"it's not the same object as ..."``. Resolving the path at *unpickle* + time works because the spawn worker imports prismatic fresh — the framework + patch is not applied there, so the module attribute is the original class. + """ + + base_module = original.__module__ + base_qualname = original.__qualname__ + + class _SpawnPicklableDataset(original): + def __init__(self, *args: Any, **kwargs: Any) -> None: + # Keep the construction arguments for __reduce__ (spawn pickling). + self.__ctor_args = args + self.__ctor_kwargs = kwargs + super().__init__(*args, **kwargs) + + def __reduce__(self): + # The TF graph cannot be pickled; send the constructor arguments and a + # *path* to the base class, then rebuild the dataset from scratch inside + # the spawn worker. + return ( + _reconstruct_rlds_dataset, + (base_module, base_qualname, self.__ctor_args, self.__ctor_kwargs), + ) + + _SpawnPicklableDataset.__turbo_physai_rlds_dataset__ = True + return _install_identity(_SpawnPicklableDataset, original) + + +def _reconstruct_rlds_dataset( + module_name: str, class_path: str, args: tuple, kwargs: dict +): + """Module-level reconstructor used by ``__reduce__`` (must be importable). + + Resolves the class from ``module_name`` / ``class_path`` **in the unpickling + process** (a spawn DataLoader worker, where prismatic is imported fresh and + the module attribute is the unpatched original class), then runs its real + constructor to rebuild the TF graph. + """ + import importlib + + module = importlib.import_module(module_name) + cls = module + for part in class_path.split("."): + cls = getattr(cls, part) + cls = _seeded_dataset_class(cls) + return cls(*args, **kwargs) + + +def _seeded_dataset_class(cls: type) -> type: + """Let the reproducibility Groups give each worker its own data stream. + + This is the only place where patched behaviour can enter a spawn worker: the + worker re-imports the model stack, so the framework's replacements are not + installed there, while this function is reached through the pickle payload. + + The hook is inert unless ``openvla.reproducibility.*`` published a run seed in + the parent process (its environment marker is inherited by the worker), i.e. + this Group keeps its exact behaviour on its own. + """ + try: + from .reproducibility import worker_seeded_dataset_class + except Exception: # pragma: no cover - a broken import must not break pickling + return cls + return worker_seeded_dataset_class(cls) + + +def rlds_dataset_spawn_wrapper( + original: type, options: Optional[Mapping[str, Any]] = None +) -> type: + """``wrap`` factory for ``RLDSDataset`` (``(original, options) -> subclass``).""" + + del options + return _spawn_picklable_dataset(original) + + +def episodic_rlds_dataset_spawn_wrapper( + original: type, options: Optional[Mapping[str, Any]] = None +) -> type: + """``wrap`` factory for ``EpisodicRLDSDataset`` (``(original, options) -> subclass``).""" + + del options + return _spawn_picklable_dataset(original) + + +# --- DataLoader spawn/worker wrapper ------------------------------------------ + +def dataloader_spawn_wrapper( + original: type, options: Optional[Mapping[str, Any]] = None +) -> type: + """``wrap`` factory for ``torch.utils.data.DataLoader``. + + Returned subclass forces ``num_workers=N`` (Group option, default 1), the + ``spawn`` multiprocessing context and — when ``N > 0`` — ``pin_memory=True`` + (Group option ``pin_memory``, default ``True``) **only when the dataset is an + RLDS dataset**; all other DataLoader constructions pass through unchanged. + + The `num_workers` option lives in this Group rather than in the upstream + DataLoader call so the item stays switchable/AB-testable from the recipe. + """ + + num_workers = _rlds_num_workers(options) + pin_memory = _rlds_pin_memory(options) + + class _SpawnWorkerDataLoader(original): + def __init__(self, dataset: Any, *args: Any, **kwargs: Any) -> None: + if _looks_like_rlds_dataset(dataset): + kwargs["num_workers"] = num_workers + if num_workers > 0: + kwargs["multiprocessing_context"] = _spawn_context() + # Pinned host memory only pays off when a worker builds the batch + # ahead of the main process; with `num_workers=0` there is nothing + # to overlap with, so leave the baseline default (False) in place. + if pin_memory: + kwargs["pin_memory"] = True + super().__init__(dataset, *args, **kwargs) + + return _install_identity(_SpawnWorkerDataLoader, original) + + +_SPAWN_CONTEXT = None + + +def _spawn_context(): + global _SPAWN_CONTEXT + if _SPAWN_CONTEXT is None: + import torch.multiprocessing as torch_mp + + _SPAWN_CONTEXT = torch_mp.get_context("spawn") + return _SPAWN_CONTEXT diff --git a/turbo_physai/optimizations/models/openvla/text_len_bucket.py b/turbo_physai/optimizations/models/openvla/text_len_bucket.py new file mode 100644 index 0000000..fc13f60 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/text_len_bucket.py @@ -0,0 +1,107 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +""" +OpenVLA fixed-length (bucketed) text padding wrapper for TurboPhysAI. +""" + +from __future__ import annotations + +import functools +import logging +from collections.abc import Mapping +from typing import Any, Optional + +import torch + +logger = logging.getLogger("openvla.data.text_len_bucket") + + +def _bucket_size(options: Optional[Mapping[str, Any]]) -> int: + options = dict(options or {}) + raw = options.get("bucket", 8) + try: + return int(raw) + except (TypeError, ValueError) as exc: + raise ValueError( + f"openvla.data.text_len_bucket: options.bucket must be an integer, got {raw!r}" + ) from exc + + +def _ignore_index() -> int: + from prismatic.util.data_utils import IGNORE_INDEX + return IGNORE_INDEX + + +def bucketed_collate_wrapper(original: Any, options: Optional[Mapping[str, Any]] = None): + """Wrapper factory ``(original, options) -> fixed-length-padding ``__call__``. + + ``original`` is ``PaddedCollatorForActionPrediction.__call__``. The returned + function calls it unchanged and then rounds the batch up to a multiple of + ``options.bucket`` tokens. + """ + # Defensive: the framework only ever passes the resolved method here. If + # something else was resolved, stay a no-op rather than breaking collation. + if not callable(original): + return original + + bucket = _bucket_size(options) + if bucket < 2: + logger.info( + "Fixed-length text padding disabled (options.bucket=%s < 2); baseline collation kept.", + bucket, + ) + return original + + # Fail fast (at apply time, before the first step) if the baseline constant + # cannot be resolved -- better than padding with a wrong ignore index. + ignore_index = _ignore_index() + + @functools.wraps(original) + def __call__(self: Any, instances: Any) -> Any: + batch = original(self, instances) + + input_ids = batch["input_ids"] + labels = batch["labels"] + batch_len = input_ids.size(1) + + # Round *up* to the next multiple, capped at `model_max_length` (which the + # baseline already truncated `input_ids` to, so the cap only ever bites + # within `bucket - 1` tokens of the limit). + bucketed_len = min( + ((batch_len + bucket - 1) // bucket) * bucket, + int(self.model_max_length), + ) + if bucketed_len <= batch_len: + return batch + + pad_len = bucketed_len - batch_len + input_ids = torch.cat( + [input_ids, input_ids.new_full((input_ids.size(0), pad_len), fill_value=self.pad_token_id)], + dim=1, + ) + labels = torch.cat( + [labels, labels.new_full((labels.size(0), pad_len), fill_value=ignore_index)], + dim=1, + ) + + batch["input_ids"] = input_ids + batch["labels"] = labels + # Same expression as the baseline computes before this point; the pads + # added above are `pad_token_id` => `attention_mask = False`. + batch["attention_mask"] = input_ids.ne(self.pad_token_id) + return batch + + __call__.__doc__ = ( + f"{original.__doc__ or ''}\n\n" + f"[openvla.data.text_len_bucket] Batch length is rounded up to a multiple of " + f"{bucket} tokens (capped at `model_max_length`); extra positions are right-padding " + f"with `pad_token_id` / label {ignore_index}." + ) + + logger.info( + "Fixed-length text padding enabled via `openvla.data.text_len_bucket` " + "(bucket=%d tokens, capped at model_max_length).", + bucket, + ) + return __call__ diff --git a/turbo_physai/optimizations/models/openvla/vision_timm.py b/turbo_physai/optimizations/models/openvla/vision_timm.py new file mode 100644 index 0000000..6dd40d3 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/vision_timm.py @@ -0,0 +1,98 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""Dynamo-safe replacement for timm's ``VisionTransformer._intermediate_layers``. + +Why this exists +--------------- +Prismatic replaces each ViT tower's ``forward`` with + + unpack_tuple(partial(featurizer.get_intermediate_layers, n={len(blocks) - 2})) + +so ``n`` is a Python ``set`` *captured outside* the compiled region. The FSDP1 +compile flow (``compile_fsdp1.py``) whole-model-``torch.compile``s each +FSDP-wrapped tower; on the first real forward Dynamo traces timm's original +line (``timm/models/vision_transformer.py``) + + take_indices = set(range(num_blocks - n, num_blocks) if isinstance(n, int) else n) + +i.e. ``set(n)`` applied to the already-existing (sourced) set. Dynamo's +builtin ``set()`` handler clones that SetVariable with +``mutation_type=ValueMutationNew()``; the clone copies the variable's +``source`` along, and ``torch/_dynamo/variables/base.py`` +(``VariableTracker.__init__``) forbids a "new" variable to carry a ``source`` +=> ``assert source is None`` raises a bare ``AssertionError`` (torch 2.7.1; +Dynamo internal limitation, not a user-code bug). + +Fix +--- +The list-based reimplementation below never runs ``set(...)`` over a sourced +set, so Dynamo traces it cleanly. It is exactly behaviour-equivalent to the +timm original for every ``n`` form the monkey-patch can pass (``int`` or any +iterable such as ``set`` / ``list`` / ``tuple`` / ``range``), so eager +numerics are unchanged. + +Wiring +------ +Declared in ``catalog.py`` as a member of the ``openvla.compile.fsdp1`` Group. +It is a ``wrap`` whose factory (``timm_intermediate_layers_wrapper``) swaps the +class attribute **only when the Group opts in via ``options.compile``** (the +vision towers are then actually ``torch.compile``d and hit the Dynamo crash). +When ``options.compile`` is false the tower stays eager, where the original +timm implementation is already fine, so the factory returns the original method +untouched (no global timm mutation). Importing this module itself has no side +effects; the class attribute is swapped only when that Group is applied with +compilation enabled. +""" + +from __future__ import annotations + +from typing import Callable, List, Sequence, Union + +import torch + + +def dynamo_safe_intermediate_layers( + self, + x: torch.Tensor, + n: Union[int, Sequence[int]] = 1, +) -> List[torch.Tensor]: + """timm ``_intermediate_layers`` with a Dynamo-safe ``take_indices``. + + Mirrors timm 0.9.16 exactly except the ``take_indices`` construction: + indices are kept in a list instead of being routed through ``set(n)`` on a + possibly-sourced set (which crashes ``torch._dynamo``, see module docstring). + """ + outputs, num_blocks = [], len(self.blocks) + if isinstance(n, int): + take_indices = list(range(num_blocks - n, num_blocks)) + else: + take_indices = list(n) # accepts set / list / tuple / range ... + + # forward pass + x = self.patch_embed(x) + x = self._pos_embed(x) + x = self.patch_drop(x) + x = self.norm_pre(x) + for i, blk in enumerate(self.blocks): + x = blk(x) + if i in take_indices: + outputs.append(x) + + return outputs + + +def timm_intermediate_layers_wrapper( + original: Callable, options +) -> Callable: + """wrap factory for ``timm...VisionTransformer._intermediate_layers``. + + The Dynamo crash this fixes only occurs when the ViT tower is actually + ``torch.compile``d. When the compile Group's ``options.compile`` is False + the tower stays eager, where the original timm implementation is already + fine, so we return the original method untouched (no global timm mutation). + """ + options = dict(options or {}) + if not options.get("compile", False): + return original + return dynamo_safe_intermediate_layers From b4cfafe37cdc8e3051ec72b4660e78d639ae264b Mon Sep 17 00:00:00 2001 From: liyh15 <73781551+xgbah@users.noreply.github.com> Date: Fri, 9 Oct 2026 09:32:34 +0800 Subject: [PATCH 2/4] feat(openvla): add packaged OpenVLA optimization configs --- .../.optimization.yaml.generation.json | 23 +++ .../models/openvla/configs/__init__.py | 4 + .../models/openvla/configs/optimization.yaml | 136 ++++++++++++++++++ .../models/openvla/configs/recipe.yaml | 83 +++++++++++ .../models/openvla/configs/runtime.yaml | 25 ++++ 5 files changed, 271 insertions(+) create mode 100644 turbo_physai/optimizations/models/openvla/configs/.optimization.yaml.generation.json create mode 100644 turbo_physai/optimizations/models/openvla/configs/__init__.py create mode 100644 turbo_physai/optimizations/models/openvla/configs/optimization.yaml create mode 100644 turbo_physai/optimizations/models/openvla/configs/recipe.yaml create mode 100644 turbo_physai/optimizations/models/openvla/configs/runtime.yaml diff --git a/turbo_physai/optimizations/models/openvla/configs/.optimization.yaml.generation.json b/turbo_physai/optimizations/models/openvla/configs/.optimization.yaml.generation.json new file mode 100644 index 0000000..9a8025e --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/configs/.optimization.yaml.generation.json @@ -0,0 +1,23 @@ +{ + "schema_version": "turbophysai/generation-record/v2", + "generator": "turbo-physai optimization generate", + "model_commit": "c8f03f48af692657d3060c19588038c7220e9af9", + "config": { + "path": "optimization.yaml", + "sha256": "e0b05c157ba5c583ba141dad1d4bbbd9e7c1b0659534318285648412253a2c3c" + }, + "inputs": [ + { + "role": "recipe", + "base": "package", + "path": "optimizations/models/openvla/configs/recipe.yaml", + "sha256": "b74db3c218df79a89fadd9dd1ca1ee11ca37ccc4731aae5977b9a458765783f8" + }, + { + "role": "catalog", + "base": "package", + "path": "optimizations/models/openvla/catalog.py", + "sha256": "899e6fc5eb45fd08b45e94669e630fd8a9f7f2fe592f963ead39a4211b671f02" + } + ] +} diff --git a/turbo_physai/optimizations/models/openvla/configs/__init__.py b/turbo_physai/optimizations/models/openvla/configs/__init__.py new file mode 100644 index 0000000..c9e2a74 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/configs/__init__.py @@ -0,0 +1,4 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +"""Packaged OpenVLA configurations.""" diff --git a/turbo_physai/optimizations/models/openvla/configs/optimization.yaml b/turbo_physai/optimizations/models/openvla/configs/optimization.yaml new file mode 100644 index 0000000..538c0e7 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/configs/optimization.yaml @@ -0,0 +1,136 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +schema_version: turbophysai/optimization-config/v1 +kind: OptimizationConfig +metadata: + id: model.openvla.base.hcu + version: 0.8.0 + description: HCU optimization recipe for official OpenVLA +model: + name: openvla +compatibility: + commits: + - c8f03f48af692657d3060c19588038c7220e9af9 +optimization_groups: +- id: openvla.bf16_support + enabled: true + options: {} + trust: + source_hashes: + prismatic.util.torch_utils.check_bloat16_supported: + - source-v1:093a3ed68b6ad038882761240a6e525f39a665318c19f077b9bfaf41d81aea9b + ast_hashes: + prismatic.util.torch_utils.check_bloat16_supported: + - ast-v1:bffb3cfaf5cec85d6e163f7b7fc1f385104af2d46b7fd851eef73578c2c26c9c +- id: openvla.fsdp.prefetch + enabled: true + options: + backward_prefetch: pre + forward_prefetch: true + limit_all_gathers: false + trust: + source_hashes: + torch.distributed.fsdp.FullyShardedDataParallel: + - source-v1:3962976c1a219799a7e6e950ba32a25da4c2d8f3f3d7a52beac24bb559a0b81a + ast_hashes: + torch.distributed.fsdp.FullyShardedDataParallel: + - ast-v1:f104543928a45a1234b0f55470af9f2223af67859c43a2a4ae23f301e381ef58 +- id: openvla.compile.fsdp1 + enabled: true + options: + compile: true + mode: default + trust: + source_hashes: + prismatic.training.strategies.fsdp.FSDPStrategy.run_setup: + - source-v1:79c3c76ef345bffe3be507b847d3f65871c8d62d8c800a8fe9e23447323da36d + prismatic.training.strategies.fsdp.FSDPStrategy.save_checkpoint: + - source-v1:f4fe324f2dda70fd17ea0c65a4015cf58dc37d715aa9b8ce1f7d38419ee8c7b4 + timm.models.vision_transformer.VisionTransformer._intermediate_layers: + - source-v1:49276f75da54527ec45b4ffcad1a54e93062fbfdf69e39d5ac9a92ea8189bbe9 + ast_hashes: + prismatic.training.strategies.fsdp.FSDPStrategy.run_setup: + - ast-v1:3334d614f6b60e577d36806093eb58cd4027e7532cd526888f9197a7794628f8 + prismatic.training.strategies.fsdp.FSDPStrategy.save_checkpoint: + - ast-v1:7db45b44ff059f737d8073e2848f9f1fef42405c49884ec64cc457e91ab011a5 + timm.models.vision_transformer.VisionTransformer._intermediate_layers: + - ast-v1:1fc1c09bd0f7b3ed6b3ced10f4835ca05e7ab26e49453c4f48d9115f17a161fa +- id: openvla.adamw.fused + enabled: true + options: {} + trust: + source_hashes: + torch.optim.AdamW: + - source-v1:887fa321b97a223e991c23a8a1ee250841bfd59cc78271039af561acaa1c8dc5 + ast_hashes: + torch.optim.AdamW: + - ast-v1:4a2ebc0bfc50c76da9c686a1d083c13231f19823863102499b0955b48453f5f5 +- id: openvla.llm.skip_fa2_unpad + enabled: true + options: {} + trust: + source_hashes: + transformers.models.llama.modeling_llama.LlamaModel._update_causal_mask: + - source-v1:61a7653e05e7e56ece80132fb619bac820f4943e46c79f32e5eab3bbcdb54552 + ast_hashes: + transformers.models.llama.modeling_llama.LlamaModel._update_causal_mask: + - ast-v1:ebea41aac0346cd6e7ce2b452c0752bb21f4eed3fe0e567fe9ed5602a4efc98e +- id: openvla.data.text_len_bucket + enabled: true + options: + bucket: 8 + trust: + source_hashes: + prismatic.util.data_utils.PaddedCollatorForActionPrediction.__call__: + - source-v1:420a81c66e38f78c4f755935363130983b35cf6a0febe6a426043e5497daa607 + ast_hashes: + prismatic.util.data_utils.PaddedCollatorForActionPrediction.__call__: + - ast-v1:99b0919ed0ec091862483cc1eced83939c13c3c652021347c4d87a44eaca04e4 +- id: openvla.dataloader.spawn + enabled: true + options: + num_workers: 2 + pin_memory: true + trust: + source_hashes: + prismatic.vla.datasets.datasets.RLDSDataset: + - source-v1:db93d5cbc56d9a4207629cc6a84fddb0d991f23814e5eb3f384e8f5cf1982516 + prismatic.vla.datasets.datasets.EpisodicRLDSDataset: + - source-v1:26bcea9656a8547e1e810d3770de4d8db92fc2af6bdd8b7e479e5a79d44d9d75 + torch.utils.data.DataLoader: + - source-v1:1e5a78c41290f043d2de3e217ab311c76ad1d294a136941c54ceb66f63514d93 + ast_hashes: + prismatic.vla.datasets.datasets.RLDSDataset: + - ast-v1:87e080d33fba03e3f0b24095802c0a2853dc55a174a51d325f8d8b971bf9ade8 + prismatic.vla.datasets.datasets.EpisodicRLDSDataset: + - ast-v1:1435c4f61646f64c321cead267b89734853af5eb186d32a8cd81624cb7a209ea + torch.utils.data.DataLoader: + - ast-v1:799d09abff6846bd4001c5250b8c87760dda1eb708dd88e18790b0232b810cc3 +- id: openvla.gc.freeze + enabled: true + options: {} + trust: + source_hashes: + prismatic.training.strategies.base_strategy.TrainingStrategy.run_vla_training: + - source-v1:0db08b36326697993fbae597026792da27067c2d9bc26800babad4e8306badfa + ast_hashes: + prismatic.training.strategies.base_strategy.TrainingStrategy.run_vla_training: + - ast-v1:a002574e805c6e58c9e10fd71f8aa8d0c539f45354d52e1d4dc58b7d0a13af25 +- id: openvla.reproducibility.data_order + enabled: false + options: + explicit_seeds: true + hold_shuffle_permutation: true + trust: {} +- id: openvla.reproducibility.determinism + enabled: false + options: + allow_tf32: false + benchmark: false + cudnn_deterministic: true + deterministic_algorithms: true + strict: false + trust: {} +optimization_modules: +- turbo_physai.optimizations.models.openvla.catalog diff --git a/turbo_physai/optimizations/models/openvla/configs/recipe.yaml b/turbo_physai/optimizations/models/openvla/configs/recipe.yaml new file mode 100644 index 0000000..fd7270c --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/configs/recipe.yaml @@ -0,0 +1,83 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +schema_version: turbophysai/optimization-config/v1 +kind: OptimizationConfig + +metadata: + id: model.openvla.base.hcu + version: "0.8.0" + description: HCU optimization recipe for official OpenVLA + +model: + name: openvla + +optimization_modules: + - turbo_physai.optimizations.models.openvla.catalog + +# NOTE: OpenVLA (prismatic) has no mmcv/mmdet3d dependency, so we do NOT extend +# `common.hcu.base` (that common collection pulls in mmcv/mmdet3d Groups whose +# targets cannot be resolved in an OpenVLA environment). This config only carries +# model-specific Groups. +extends: [] + +compatibility: {} + +# Add validated model-specific Group IDs here. +optimization_groups: + - id: openvla.bf16_support + enabled: true + + # 放在 compile 组之前:该组 patch 的是 FSDP1 类本身,compile 组的 run_setup 在运行时动态取这个类来构造 FSDP,两者因此可以叠加。 + - id: openvla.fsdp.prefetch + enabled: true + options: + limit_all_gathers: false + forward_prefetch: true + backward_prefetch: pre + + - id: openvla.compile.fsdp1 + enabled: true + options: + compile: true + mode: default + + - id: openvla.adamw.fused + enabled: true + + - id: openvla.llm.skip_fa2_unpad + enabled: true + + - id: openvla.data.text_len_bucket + enabled: true + options: + bucket: 8 + + - id: openvla.dataloader.spawn + enabled: true + options: + num_workers: 2 + pin_memory: true + + - id: openvla.gc.freeze + enabled: true + + # 可复现性(数据流):把 run seed 打进 TF 全局 RNG 并显式送给 RLDS 管线,固定 TFDS + # 文件级 shuffle、mixture 采样与 frame shuffle。与 openvla.dataloader.spawn 一起开启时, + # 每个 worker 还会拿到互不相同且可复现的数据流。 + - id: openvla.reproducibility.data_order + enabled: false + options: + explicit_seeds: true + hold_shuffle_permutation: true + + # 可复现性(算法选择):固定 cudnn/MIOpen 的自动调优与 TF32、启用确定性算子,并在模型 + # 构造后重新落实 run seed。只固定 kernel、不固定样本 —— 逐位复现需与 data_order 同开。 + - id: openvla.reproducibility.determinism + enabled: false + options: + cudnn_deterministic: true + benchmark: false + allow_tf32: false + deterministic_algorithms: true + strict: false diff --git a/turbo_physai/optimizations/models/openvla/configs/runtime.yaml b/turbo_physai/optimizations/models/openvla/configs/runtime.yaml new file mode 100644 index 0000000..58a5f20 --- /dev/null +++ b/turbo_physai/optimizations/models/openvla/configs/runtime.yaml @@ -0,0 +1,25 @@ +# Copyright 2026 Hygon Information Technology Co., Ltd. +# SPDX-License-Identifier: BSD-3-Clause + +schema_version: turbophysai/runtime-config/v1 +kind: RuntimeConfig + +# Machine-independent defaults for the OpenVLA training recipe. +# Network interface, master address/port, CPU placement and host-specific +# library paths stay deployer overrides because they differ between nodes +# and schedulers. +environment: + set: + DISABLE_ADDMM_CUDA_LT: "1" + MIOPEN_PRECISION_FP32_FP32_FP32_TF32_FP32: "1" + PYTORCH_HIP_ALLOC_CONF: "roundup_power2_divisions:16" + unset: [] + +# Not set here on purpose: ROCBLAS_TENSILE_LIBPATH. The rocBLAS tuning library +# is generated per node by `rocblas-tensile` for that node's device and driver, +# so it cannot ship as a packaged default. Deployers who have tuned a node +# override it on the command line: +# turbo-physai run --set ROCBLAS_TENSILE_LIBPATH= ... + +process: + numa: true From ea9ca27e902cf23d89fa40f0cc7aa38c74bc01ad Mon Sep 17 00:00:00 2001 From: liyh15 <73781551+xgbah@users.noreply.github.com> Date: Fri, 9 Oct 2026 10:52:47 +0800 Subject: [PATCH 3/4] docs(openvla): add OpenVLA application guide Add `model_examples/OpenVLA/README.md`, required by CONTRIBUTING for a newly supported model: model overview, the upstream OpenVLA commit the optimizations were validated against (`c8f03f4`), source / BridgeData V2 / base-VLM preparation, and the `turbo-physai run` command for single-node 8-accelerator full fine-tuning. --- model_examples/OpenVLA/README.md | 99 ++++++++++++++++++++++++++++++++ 1 file changed, 99 insertions(+) create mode 100644 model_examples/OpenVLA/README.md diff --git a/model_examples/OpenVLA/README.md b/model_examples/OpenVLA/README.md new file mode 100644 index 0000000..d67f033 --- /dev/null +++ b/model_examples/OpenVLA/README.md @@ -0,0 +1,99 @@ +# OpenVLA 应用说明 + +本文说明如何在产品镜像中,基于官方 OpenVLA 仓库应用 TurboPhysAI 的 HCU 优化。 + +## 1. 模型简介 + +[OpenVLA](https://github.com/openvla/openvla) 是 7B 规模的开源视觉-语言-动作(VLA)模型,基于 Prismatic VLM(DINOv2 + SigLIP 视觉编码器与 Llama-2 7B 语言模型)构建,把机器人操作观测和语言指令映射为离散化的 7 自由度动作 token,支持在 BridgeData V2 等操作数据集上进行全参数微调。 + +训练使用 PyTorch FSDP1 分片,入口脚本为 `vla-scripts/train.py`。TurboPhysAI 的 OpenVLA 优化覆盖该训练栈的算子、通信、数据与可复现性环节,全部通过随包 OptimizationConfig 和 RuntimeConfig 交付,不修改模型源码。 + +## 2. 优化接入基线 + +TurboPhysAI 的 OpenVLA 优化基于官方仓库 commit `c8f03f48af692657d3060c19588038c7220e9af9` 接入。优化接入基线的使用建议见[模型支持清单](../../docs/zh/models/support_list.md)。 + +## 3. 准备模型源码 + +```bash +cd /workspace +mkdir -p model +git clone https://github.com/openvla/openvla.git model/OpenVLA +cd model/OpenVLA +git checkout --detach c8f03f48af692657d3060c19588038c7220e9af9 +``` + +也可以将准备好的官方 OpenVLA 仓库放入宿主机工作目录的 `model/OpenVLA`。下文命令均在容器内的 `/workspace/model/OpenVLA` 执行。产品镜像提供 TurboPhysAI 与 HCU 运行环境,不包含模型源码;OpenVLA 训练依赖(`torch`、`transformers`、`timm`、`flash-attn`、RLDS/dlimp、`draccus` 等)按[官方安装说明](https://github.com/openvla/openvla#installation)在该仓库的训练环境中准备。 + +## 4. 准备 BridgeData V2 数据 + +OpenVLA 训练使用 RLDS 格式的数据。本文使用 BridgeData V2,从官方地址下载(约 124 GB): + +```bash +cd + +wget -r -nH --cut-dirs=4 --reject="index.html*" \ + https://rail.eecs.berkeley.edu/datasets/bridge_release/data/tfds/bridge_dataset/ + +# 目录名必须是 bridge_orig,否则启动后会因找不到数据集而报错 +mv bridge_dataset bridge_orig +``` + +数据根目录下的结构应为: + +```text +/ +└── bridge_orig/ + └── 1.0.0/ + ├── dataset_info.json + └── bridge_orig-train.tfrecord-* +``` + +启动训练时通过 `--data_root_dir` 传入该数据根目录。TurboPhysAI 不下载、转换或重新分发数据集。 + +## 5. 准备基座 VLM 权重 + +`--vla.type prism-dinosiglip-224px+mx-bridge` 的 `base_vlm` 为 `prism-dinosiglip-224px+7b`,训练脚本在未指定 `--pretrained_checkpoint` 时从 Hugging Face Hub 拉取该 Prismatic VLM。按官方说明在仓库根目录准备 Hugging Face token: + +```bash +cd /workspace/model/OpenVLA + +# 将 Hugging Face token(形如 hf_...)写入仓库根目录的 .hf_token +printf '%s\n' 'hf_xxxxxxxxxxxxxxxxxxxxxxxx' > .hf_token +``` + +离线环境可以先在有网络的机器上填充 Hugging Face 缓存,再把缓存目录挂载进容器并通过 `HF_HOME` 指向它。也可以改为加载本地检查点,此时 `--pretrained_checkpoint` 取代上面的 Hub 下载路径(`--vla.type` 仍需保留,用于选择训练配置): + +```bash +--pretrained_checkpoint +``` + +## 6. 通过 turbo-physai run 启动训练 + +`turbo-physai run` 自动加载随包交付的 [OptimizationConfig](../../docs/zh/user_guide/optimization_config.md) 和 [RuntimeConfig](../../docs/zh/user_guide/runtime_config.md),无需修改模型源码即可应用 OpenVLA 优化。 + +### 6.1 单机八卡训练 + +```bash +source /opt/conda/etc/profile.d/conda.sh +conda activate +cd /workspace/model/OpenVLA + +turbo-physai run \ + --model openvla \ + --log-report \ + torchrun --standalone --nnodes 1 --nproc-per-node 8 vla-scripts/train.py \ + --vla.type prism-dinosiglip-224px+mx-bridge \ + --vla.train_strategy fsdp-shard-grad-op \ + --vla.expected_world_size 8 \ + --vla.global_batch_size 256 \ + --vla.per_device_batch_size 32 \ + --vla.enable_gradient_checkpointing true \ + --vla.reduce_in_full_precision true \ + --vla.max_steps 1000 \ + --data_root_dir \ + --run_root_dir <训练输出目录> \ + --run_id_note openvla_8card \ + --trackers '[jsonl]' +``` + +需要使用自定义交付配置时,通过 `--optimization-config` 和 `--runtime-config` 指定对应文件。 From e7aaad0202cf2fb7ccf8082765bfba2057e00517 Mon Sep 17 00:00:00 2001 From: liyh15 <73781551+xgbah@users.noreply.github.com> Date: Fri, 9 Oct 2026 10:52:47 +0800 Subject: [PATCH 4/4] docs(models): list OpenVLA in the supported-model table Record the OpenVLA optimization baseline (`c8f03f4`) and link the new application guide from `docs/zh/models/support_list.md`. --- docs/zh/models/support_list.md | 1 + 1 file changed, 1 insertion(+) diff --git a/docs/zh/models/support_list.md b/docs/zh/models/support_list.md index 4df19cf..a2dd643 100644 --- a/docs/zh/models/support_list.md +++ b/docs/zh/models/support_list.md @@ -12,5 +12,6 @@ TurboPhysAI 不强制模型仓库停留在优化接入基线([可以在优化 | :-------: | :------------: | :------: | :------------------------------------------: | :---------------------------------------------------: | | BEVFormer | BEVFormer-base | R101-DCN | `66b65f3a1f58caf0507cb2a971b9c0e7f842376c` | [BEVFormer](../../../model_examples/BEVFormer/README.md) | | BEVFusion | BEVFusion | — | `326653dc06e0938edf1aae7d01efcd158ba83de5` | [BEVFusion](../../../model_examples/BEVFusion/README.md) | +| OpenVLA | OpenVLA-7B | DINOv2+SigLIP | `c8f03f48af692657d3060c19588038c7220e9af9` | [OpenVLA](../../../model_examples/OpenVLA/README.md) | **自动驾驶**、**具身智能**和**世界模型**等方向的模型适配工作持续推进中,请关注后续更新。