mirror of
https://github.com/huggingface/lerobot.git
synced 2026-08-08 17:39:44 +00:00
refactor(g05): keep imports at module scope
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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]),
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user