refactor(g05): keep imports at module scope

This commit is contained in:
Pepijn
2026-07-30 10:42:48 +02:00
parent 92c2c6cffd
commit 8aec9a1bb7
4 changed files with 19 additions and 48 deletions
+13 -21
View File
@@ -16,6 +16,8 @@
from __future__ import annotations
import itertools
import json
import math
import time
from collections.abc import Mapping
@@ -25,6 +27,17 @@ from typing import Any
import torch
import torch.nn.functional as functional
from torch import Tensor, nn
from transformers import DynamicCache
from transformers.models.qwen3_5.configuration_qwen3_5 import Qwen3_5TextConfig, Qwen3_5VisionConfig
from transformers.models.qwen3_5.modeling_qwen3_5 import (
Qwen3_5Attention,
Qwen3_5DecoderLayer,
Qwen3_5MLP,
Qwen3_5RMSNorm,
Qwen3_5TextRotaryEmbedding,
Qwen3_5VisionModel,
apply_rotary_pos_emb_vision,
)
from lerobot.policies.pi_gemma import PiGemmaRMSNorm
from lerobot.utils.constants import ACTION
@@ -38,8 +51,6 @@ G05_RUNTIME_PREDICT_COT = "g05_runtime_predict_cot"
def _qwen_text_config(values: Mapping[str, Any], *, vocab_size: int | None = None):
"""Translate the serialized G0.5 Qwen config into a Transformers config."""
from transformers.models.qwen3_5.configuration_qwen3_5 import Qwen3_5TextConfig
return Qwen3_5TextConfig(
vocab_size=int(vocab_size if vocab_size is not None else values.get("vocab_size", 1)),
hidden_size=int(values["hidden_size"]),
@@ -67,8 +78,6 @@ def _qwen_text_config(values: Mapping[str, Any], *, vocab_size: int | None = Non
def _qwen_vision_config(values: Mapping[str, Any]):
"""Translate the serialized G0.5 vision config into Transformers."""
from transformers.models.qwen3_5.configuration_qwen3_5 import Qwen3_5VisionConfig
config = Qwen3_5VisionConfig(
depth=int(values["depth"]),
hidden_size=int(values["hidden_size"]),
@@ -112,11 +121,6 @@ class G05QwenTextModel(nn.Module):
def __init__(self, values: Mapping[str, Any], *, vocab_size: int) -> None:
super().__init__()
from transformers.models.qwen3_5.modeling_qwen3_5 import (
Qwen3_5DecoderLayer,
Qwen3_5RMSNorm,
Qwen3_5TextRotaryEmbedding,
)
self.config = _qwen_text_config(values, vocab_size=vocab_size)
self.input_proj = nn.Embedding(vocab_size, self.config.hidden_size, self.config.pad_token_id)
@@ -144,8 +148,6 @@ class G05QwenTextModel(nn.Module):
position_ids: Tensor,
cache=None,
) -> tuple[Tensor, Any]:
from transformers import DynamicCache
if cache is None:
cache = DynamicCache(config=self.config)
position_embeddings = self.rotary_emb(inputs_embeds, position_ids)
@@ -172,7 +174,6 @@ class G05ActionDecoderLayer(nn.Module):
def __init__(self, config, layer_idx: int) -> None:
super().__init__()
from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5Attention, Qwen3_5MLP
if config.layer_types[layer_idx] != "full_attention":
raise ValueError("The released G0.5 action expert requires full-attention layers.")
@@ -227,7 +228,6 @@ class G05ActionExpert(nn.Module):
def __init__(self, values: Mapping[str, Any]) -> None:
super().__init__()
from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5TextRotaryEmbedding
self.config = _qwen_text_config(values)
input_dim = int(values["input_dim"])
@@ -297,7 +297,6 @@ class G05NativeModel(nn.Module):
def __init__(self, model_config: Mapping[str, Any], *, vocab_size: int) -> None:
super().__init__()
from transformers.models.qwen3_5.modeling_qwen3_5 import Qwen3_5VisionModel
self.vision_tower = Qwen3_5VisionModel(_qwen_vision_config(model_config["vision"]))
self.vlm = G05QwenTextModel(model_config["vlm"], vocab_size=vocab_size)
@@ -463,7 +462,6 @@ class G05NativeBackend(nn.Module):
tokenizer_config = processor_path / "tokenizer_config.json"
if not tokenizer_config.is_file():
raise FileNotFoundError(f"G0.5 tokenizer config not found: {tokenizer_config}")
import json
tokenizer_metadata = json.loads(tokenizer_config.read_text())
added = tokenizer_metadata.get("added_tokens_decoder") or {}
@@ -513,8 +511,6 @@ class G05NativeBackend(nn.Module):
temporal_pe: Tensor,
temporal_mask: Tensor,
) -> Tensor:
from transformers.models.qwen3_5.modeling_qwen3_5 import apply_rotary_pos_emb_vision
total, hidden_size = hidden_states.shape
num_heads = block.attn.num_heads
head_dim = block.attn.head_dim
@@ -703,8 +699,6 @@ class G05NativeBackend(nn.Module):
return embeddings
def _mrope_positions(self, token_types: Tensor) -> Tensor:
import itertools
batch_size, sequence_length = token_types.shape
positions = torch.zeros(
3,
@@ -894,8 +888,6 @@ class G05NativeBackend(nn.Module):
return generated_ids, cache, last_hidden, token_types, positions
def _action_cache(self, vlm_cache, prefix_length: int, *, repeats: int = 1):
from transformers import DynamicCache
cache = DynamicCache(config=self.model.action_expert.config)
layer_types = self.model.vlm.config.layer_types
for layer_index, layer_type in enumerate(layer_types):
+1 -2
View File
@@ -23,6 +23,7 @@ from typing import Any
import torch
from torch import Tensor
from transformers import AutoTokenizer
IGNORE_INDEX = -100
@@ -68,8 +69,6 @@ class G05Tokenizer:
_PLACEHOLDER = re.compile(r"<([^<>|]+)>")
def __init__(self, processor_path: str | Path, model_config: dict[str, Any]) -> None:
from transformers import AutoTokenizer
self.processor_path = Path(processor_path)
self.tokenizer = AutoTokenizer.from_pretrained(
self.processor_path,
+3 -6
View File
@@ -17,6 +17,8 @@ from typing import Any
import torch
import torchvision.transforms.functional as vision_functional
from lerobot.configs import recipe as recipe_module
from lerobot.configs.recipe import TrainingRecipe
from lerobot.configs.types import FeatureType, NormalizationMode, PipelineFeatureType, PolicyFeature
from lerobot.processor import (
AbsoluteActionsProcessorStep,
@@ -38,6 +40,7 @@ from lerobot.processor.converters import (
transition_to_policy_action,
)
from lerobot.processor.relative_action_processor import to_relative_actions
from lerobot.processor.render_messages_processor import RenderMessagesStep
from lerobot.types import EnvTransition, TransitionKey
from lerobot.utils.constants import (
ACTION,
@@ -52,12 +55,8 @@ from .configuration_g05 import G05_EMBODIMENT_MAPPINGS, G05_POLICY_PARTS, G05Con
def _load_recipe(path_str: str) -> Any:
"""Load an absolute recipe path or one relative to ``lerobot/configs``."""
from lerobot.configs.recipe import TrainingRecipe
path = Path(path_str)
if not path.is_absolute() and not path.exists():
from lerobot.configs import recipe as recipe_module
candidate = Path(recipe_module.__file__).resolve().parent / path
if candidate.exists():
path = candidate
@@ -740,8 +739,6 @@ def make_g05_pre_post_processors(
)
steps: list[ProcessorStep] = [RenameObservationsProcessorStep(rename_map={})]
if config.recipe_path:
from lerobot.processor.render_messages_processor import RenderMessagesStep
steps.extend(
[
G05BBoxImageSizeStep(camera_key=config.cot_bbox_camera or config.camera_order[0]),
+2 -19
View File
@@ -16,8 +16,10 @@ from types import SimpleNamespace
import pytest
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
from lerobot.runtime import LanguageConditionedRuntime, RuntimeState
from lerobot.runtime.adapter import GenerationConfig
from lerobot.runtime.registry import get_language_adapter_factory
class FakeG05Policy:
@@ -36,15 +38,10 @@ class FakeG05Policy:
def test_registry_lazily_resolves_g05_adapter():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
from lerobot.runtime.registry import get_language_adapter_factory
assert get_language_adapter_factory("g05") is G05PolicyAdapter
def test_system1_passes_exact_runtime_task_without_mutating_observation():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
policy = FakeG05Policy()
adapter = G05PolicyAdapter(policy)
original = {"task": "stale generated subtask", "observation.state": "state"}
@@ -58,8 +55,6 @@ def test_system1_passes_exact_runtime_task_without_mutating_observation():
def test_system2_surfaces_same_pass_cot_and_action():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
class ReasoningPolicy(FakeG05Policy):
def __init__(self):
super().__init__(predict_cot=True, continuous_action=True)
@@ -90,8 +85,6 @@ def test_system2_surfaces_same_pass_cot_and_action():
def test_system2_accepts_batch_safe_tuple_metadata():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
class ReasoningPolicy(FakeG05Policy):
def __init__(self):
super().__init__(predict_cot=True)
@@ -108,8 +101,6 @@ def test_system2_accepts_batch_safe_tuple_metadata():
def test_system2_reasoning_does_not_invalidate_same_pass_action_chunk():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
class ReasoningPolicy(FakeG05Policy):
def __init__(self):
super().__init__(predict_cot=True)
@@ -133,30 +124,22 @@ def test_system2_reasoning_does_not_invalidate_same_pass_action_chunk():
def test_system2_rejects_checkpoint_without_predict_cot():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
with pytest.raises(ValueError, match="predict_cot=True"):
G05PolicyAdapter(FakeG05Policy(predict_cot=False), system_mode="system2")
def test_system1_rejects_checkpoint_without_an_action_head():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
with pytest.raises(ValueError, match="both discrete_action=False and continuous_action=False"):
G05PolicyAdapter(FakeG05Policy(predict_cot=True, discrete_action=False, continuous_action=False))
def test_system2_requires_structured_single_pass_hook():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
adapter = G05PolicyAdapter(FakeG05Policy(predict_cot=True))
with pytest.raises(RuntimeError, match="predict_action_chunk_with_runtime"):
adapter.select_action({}, RuntimeState(task="pick"))
def test_direct_subtask_selects_system1_on_system2_checkpoint():
from lerobot.policies.g05.inference.g05_adapter import G05PolicyAdapter
class SwitchablePolicy(FakeG05Policy):
def __init__(self):
super().__init__(predict_cot=True, continuous_action=True)