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