refactor(xvla): reuse native Florence2 components (#4089)

This commit is contained in:
Steven Palma
2026-07-20 19:19:41 +02:00
committed by GitHub
parent ddc2aa7a27
commit 1bb9933215
4 changed files with 167 additions and 3185 deletions
@@ -1,355 +0,0 @@
# Copyright 2024 Microsoft and the HuggingFace Inc. team. All rights reserved.
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
import warnings
from transformers.configuration_utils import PretrainedConfig
from transformers.utils import logging
""" Florence-2 configuration"""
logger = logging.get_logger(__name__)
class Florence2VisionConfig(PretrainedConfig):
r"""
This is the configuration class to store the configuration of a [`Florence2VisionModel`]. It is used to instantiate a Florence2VisionModel
according to the specified arguments, defining the model architecture. Instantiating a configuration with the
defaults will yield a similar configuration to that of the Florence2VisionModel architecture.
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
documentation from [`PretrainedConfig`] for more information.
Args:
drop_path_rate (`float`, *optional*, defaults to 0.1):
The dropout rate of the drop path layer.
patch_size (`List[int]`, *optional*, defaults to [7, 3, 3, 3]):
The patch size of the image.
patch_stride (`List[int]`, *optional*, defaults to [4, 2, 2, 2]):
The patch stride of the image.
patch_padding (`List[int]`, *optional*, defaults to [3, 1, 1, 1]):
The patch padding of the image.
patch_prenorm (`List[bool]`, *optional*, defaults to [false, true, true, true]):
Whether to apply layer normalization before the patch embedding layer.
enable_checkpoint (`bool`, *optional*, defaults to False):
Whether to enable checkpointing.
dim_embed (`List[int]`, *optional*, defaults to [256, 512, 1024, 2048]):
The dimension of the embedding layer.
num_heads (`List[int]`, *optional*, defaults to [8, 16, 32, 64]):
The number of attention heads.
num_groups (`List[int]`, *optional*, defaults to [8, 16, 32, 64]):
The number of groups.
depths (`List[int]`, *optional*, defaults to [1, 1, 9, 1]):
The depth of the model.
window_size (`int`, *optional*, defaults to 12):
The window size of the model.
projection_dim (`int`, *optional*, defaults to 1024):
The dimension of the projection layer.
visual_temporal_embedding (`dict`, *optional*):
The configuration of the visual temporal embedding.
image_pos_embed (`dict`, *optional*):
The configuration of the image position embedding.
image_feature_source (`List[str]`, *optional*, defaults to ["spatial_avg_pool", "temporal_avg_pool"]):
The source of the image feature.
Example:
```python
>>> from transformers import Florence2VisionConfig, Florence2VisionModel
>>> # Initializing a Florence2 Vision style configuration
>>> configuration = Florence2VisionConfig()
>>> # Initializing a model (with random weights)
>>> model = Florence2VisionModel(configuration)
>>> # Accessing the model configuration
>>> configuration = model.config
```"""
model_type = "davit"
keys_to_ignore_at_inference = ["past_key_values"]
def __init__(
self,
drop_path_rate=0.1,
patch_size=None,
patch_stride=None,
patch_padding=None,
patch_prenorm=None,
enable_checkpoint=False,
dim_embed=None,
num_heads=None,
num_groups=None,
depths=None,
window_size=12,
projection_dim=1024,
visual_temporal_embedding=None,
image_pos_embed=None,
image_feature_source=None,
**kwargs,
):
self.drop_path_rate = drop_path_rate
self.patch_size = patch_size if patch_size is not None else [7, 3, 3, 3]
self.patch_stride = patch_stride if patch_stride is not None else [4, 2, 2, 2]
self.patch_padding = patch_padding if patch_padding is not None else [3, 1, 1, 1]
self.patch_prenorm = patch_prenorm if patch_prenorm is not None else [False, True, True, True]
self.enable_checkpoint = enable_checkpoint
self.dim_embed = dim_embed if dim_embed is not None else [256, 512, 1024, 2048]
self.num_heads = num_heads if num_heads is not None else [8, 16, 32, 64]
self.num_groups = num_groups if num_groups is not None else [8, 16, 32, 64]
self.depths = depths if depths is not None else [1, 1, 9, 1]
self.window_size = window_size
self.projection_dim = projection_dim
if visual_temporal_embedding is None:
visual_temporal_embedding = {
"type": "COSINE",
"max_temporal_embeddings": 100,
}
self.visual_temporal_embedding = visual_temporal_embedding
if image_pos_embed is None:
image_pos_embed = {
"type": "learned_abs_2d",
"max_pos_embeddings": 1000,
}
self.image_pos_embed = image_pos_embed
self.image_feature_source = (
image_feature_source
if image_feature_source is not None
else ["spatial_avg_pool", "temporal_avg_pool"]
)
super().__init__(**kwargs)
class Florence2LanguageConfig(PretrainedConfig):
r"""
This is the configuration class to store the configuration of a [`Florence2LanguagePreTrainedModel`]. It is used to instantiate a BART
model according to the specified arguments, defining the model architecture. Instantiating a configuration with the
defaults will yield a similar configuration to that of the BART
[facebook/bart-large](https://huggingface.co/facebook/bart-large) architecture.
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
documentation from [`PretrainedConfig`] for more information.
Args:
vocab_size (`int`, *optional*, defaults to 51289):
Vocabulary size of the Florence2Language model. Defines the number of different tokens that can be represented by the
`inputs_ids` passed when calling [`Florence2LanguageModel`].
d_model (`int`, *optional*, defaults to 1024):
Dimensionality of the layers and the pooler layer.
encoder_layers (`int`, *optional*, defaults to 12):
Number of encoder layers.
decoder_layers (`int`, *optional*, defaults to 12):
Number of decoder layers.
encoder_attention_heads (`int`, *optional*, defaults to 16):
Number of attention heads for each attention layer in the Transformer encoder.
decoder_attention_heads (`int`, *optional*, defaults to 16):
Number of attention heads for each attention layer in the Transformer decoder.
decoder_ffn_dim (`int`, *optional*, defaults to 4096):
Dimensionality of the "intermediate" (often named feed-forward) layer in decoder.
encoder_ffn_dim (`int`, *optional*, defaults to 4096):
Dimensionality of the "intermediate" (often named feed-forward) layer in decoder.
activation_function (`str` or `function`, *optional*, defaults to `"gelu"`):
The non-linear activation function (function or string) in the encoder and pooler. If string, `"gelu"`,
`"relu"`, `"silu"` and `"gelu_new"` are supported.
dropout (`float`, *optional*, defaults to 0.1):
The dropout probability for all fully connected layers in the embeddings, encoder, and pooler.
attention_dropout (`float`, *optional*, defaults to 0.0):
The dropout ratio for the attention probabilities.
activation_dropout (`float`, *optional*, defaults to 0.0):
The dropout ratio for activations inside the fully connected layer.
classifier_dropout (`float`, *optional*, defaults to 0.0):
The dropout ratio for classifier.
max_position_embeddings (`int`, *optional*, defaults to 1024):
The maximum sequence length that this model might ever be used with. Typically set this to something large
just in case (e.g., 512 or 1024 or 2048).
init_std (`float`, *optional*, defaults to 0.02):
The standard deviation of the truncated_normal_initializer for initializing all weight matrices.
encoder_layerdrop (`float`, *optional*, defaults to 0.0):
The LayerDrop probability for the encoder. See the [LayerDrop paper](see https://arxiv.org/abs/1909.11556)
for more details.
decoder_layerdrop (`float`, *optional*, defaults to 0.0):
The LayerDrop probability for the decoder. See the [LayerDrop paper](see https://arxiv.org/abs/1909.11556)
for more details.
scale_embedding (`bool`, *optional*, defaults to `False`):
Scale embeddings by diving by sqrt(d_model).
use_cache (`bool`, *optional*, defaults to `True`):
Whether or not the model should return the last key/values attentions (not used by all models).
num_labels (`int`, *optional*, defaults to 3):
The number of labels to use in [`Florence2LanguageForSequenceClassification`].
forced_eos_token_id (`int`, *optional*, defaults to 2):
The id of the token to force as the last generated token when `max_length` is reached. Usually set to
`eos_token_id`.
Example:
```python
>>> from transformers import Florence2LanguageConfig, Florence2LanguageModel
>>> # Initializing a Florence2 Language style configuration
>>> configuration = Florence2LanguageConfig()
>>> # Initializing a model (with random weights)
>>> model = Florence2LanguageModel(configuration)
>>> # Accessing the model configuration
>>> configuration = model.config
```"""
model_type = "florence2_language"
keys_to_ignore_at_inference = ["past_key_values"]
attribute_map = {"num_attention_heads": "encoder_attention_heads", "hidden_size": "d_model"}
def __init__(
self,
vocab_size=51289,
max_position_embeddings=1024,
encoder_layers=12,
encoder_ffn_dim=4096,
encoder_attention_heads=16,
decoder_layers=12,
decoder_ffn_dim=4096,
decoder_attention_heads=16,
encoder_layerdrop=0.0,
decoder_layerdrop=0.0,
activation_function="gelu",
d_model=1024,
dropout=0.1,
attention_dropout=0.0,
activation_dropout=0.0,
init_std=0.02,
classifier_dropout=0.0,
scale_embedding=False,
use_cache=True,
num_labels=3,
pad_token_id=1,
bos_token_id=0,
eos_token_id=2,
is_encoder_decoder=True,
decoder_start_token_id=2,
forced_eos_token_id=2,
**kwargs,
):
self.vocab_size = vocab_size
self.max_position_embeddings = max_position_embeddings
self.d_model = d_model
self.encoder_ffn_dim = encoder_ffn_dim
self.encoder_layers = encoder_layers
self.encoder_attention_heads = encoder_attention_heads
self.decoder_ffn_dim = decoder_ffn_dim
self.decoder_layers = decoder_layers
self.decoder_attention_heads = decoder_attention_heads
self.dropout = dropout
self.attention_dropout = attention_dropout
self.activation_dropout = activation_dropout
self.activation_function = activation_function
self.init_std = init_std
self.encoder_layerdrop = encoder_layerdrop
self.decoder_layerdrop = decoder_layerdrop
self.classifier_dropout = classifier_dropout
self.use_cache = use_cache
self.num_hidden_layers = encoder_layers
self.scale_embedding = scale_embedding # scale factor will be sqrt(d_model) if True
super().__init__(
num_labels=num_labels,
pad_token_id=pad_token_id,
bos_token_id=bos_token_id,
eos_token_id=eos_token_id,
is_encoder_decoder=is_encoder_decoder,
decoder_start_token_id=decoder_start_token_id,
forced_eos_token_id=forced_eos_token_id,
**kwargs,
)
# ensure backward compatibility for BART CNN models
if not hasattr(self, "forced_bos_token_id"):
self.forced_bos_token_id = None
if self.forced_bos_token_id is None and kwargs.get("force_bos_token_to_be_generated", False):
self.forced_bos_token_id = self.bos_token_id
warnings.warn(
f"Please make sure the config includes `forced_bos_token_id={self.bos_token_id}` in future versions. "
"The config can simply be saved and uploaded again to be fixed.",
stacklevel=2,
)
class Florence2Config(PretrainedConfig):
r"""
This is the configuration class to store the configuration of a [`Florence2ForConditionalGeneration`]. It is used to instantiate an
Florence-2 model according to the specified arguments, defining the model architecture.
Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the
documentation from [`PretrainedConfig`] for more information.
Args:
vision_config (`Florence2VisionConfig`, *optional*):
Custom vision config or dict
text_config (`Union[AutoConfig, dict]`, *optional*):
The config object of the text backbone.
ignore_index (`int`, *optional*, defaults to -100):
The ignore index for the loss function.
vocab_size (`int`, *optional*, defaults to 51289):
Vocabulary size of the Florence2model. Defines the number of different tokens that can be represented by the
`inputs_ids` passed when calling [`~Florence2ForConditionalGeneration`]
projection_dim (`int`, *optional*, defaults to 1024):
Dimension of the multimodal projection space.
Example:
```python
>>> from transformers import Florence2ForConditionalGeneration, Florence2Config, CLIPVisionConfig, BartConfig
>>> # Initializing a clip-like vision config
>>> vision_config = CLIPVisionConfig()
>>> # Initializing a Bart config
>>> text_config = BartConfig()
>>> # Initializing a Florence-2 configuration
>>> configuration = Florence2Config(vision_config, text_config)
>>> # Initializing a model from the florence-2 configuration
>>> model = Florence2ForConditionalGeneration(configuration)
>>> # Accessing the model configuration
>>> configuration = model.config
```"""
model_type = "florence2"
is_composition = False
def __init__(
self,
vision_config=None,
text_config=None,
ignore_index=-100,
vocab_size=51289,
projection_dim=1024,
**kwargs,
):
self.ignore_index = ignore_index
self.vocab_size = vocab_size
self.projection_dim = projection_dim
if vision_config is not None:
vision_config = Florence2VisionConfig(**vision_config)
self.vision_config = vision_config
self.text_config = text_config
if text_config is not None:
self.text_config = Florence2LanguageConfig(**text_config)
super().__init__(**kwargs)
@@ -29,11 +29,50 @@ from lerobot.utils.constants import OBS_IMAGES
from lerobot.utils.import_utils import _transformers_available from lerobot.utils.import_utils import _transformers_available
if TYPE_CHECKING or _transformers_available: if TYPE_CHECKING or _transformers_available:
from .configuration_florence2 import Florence2Config from transformers import Florence2Config
else: else:
Florence2Config = None Florence2Config = None
def _translate_vision_config(vision_config: dict[str, Any]) -> dict[str, Any]:
"""Translate a vision config from the original Microsoft remote-code Florence-2 format
(used by existing XVLA checkpoints) to the native ``transformers`` format.
Configs already in the native format pass through unchanged.
"""
vision = dict(vision_config)
model_type = vision.pop("model_type", None)
if model_type not in (None, "davit", "florence_vision"):
raise ValueError(f"Unsupported Florence-2 vision backbone: {model_type!r}")
vision.pop("enable_checkpoint", None)
image_pos_embed = vision.pop("image_pos_embed", None)
if image_pos_embed is not None:
if image_pos_embed.get("type") != "learned_abs_2d":
raise ValueError(f"Unsupported image_pos_embed type: {image_pos_embed.get('type')!r}")
vision["max_position_embeddings"] = image_pos_embed["max_pos_embeddings"]
visual_temporal_embedding = vision.pop("visual_temporal_embedding", None)
if visual_temporal_embedding is not None:
if visual_temporal_embedding.get("type") != "COSINE":
raise ValueError(
f"Unsupported visual_temporal_embedding type: {visual_temporal_embedding.get('type')!r}"
)
vision["max_temporal_embeddings"] = visual_temporal_embedding["max_temporal_embeddings"]
image_feature_source = vision.pop("image_feature_source", None)
if image_feature_source is not None and list(image_feature_source) != [
"spatial_avg_pool",
"temporal_avg_pool",
]:
# the native Florence2MultiModalProjector hardcodes this feature combination
raise ValueError(f"Unsupported image_feature_source: {image_feature_source!r}")
if "dim_embed" in vision:
vision["embed_dim"] = vision.pop("dim_embed")
return vision
@PreTrainedConfig.register_subclass("xvla") @PreTrainedConfig.register_subclass("xvla")
@dataclass @dataclass
class XVLAConfig(PreTrainedConfig): class XVLAConfig(PreTrainedConfig):
@@ -128,16 +167,41 @@ class XVLAConfig(PreTrainedConfig):
def get_florence_config(self) -> Florence2Config: def get_florence_config(self) -> Florence2Config:
""" """
Build (and cache) the Florence2 transformer config that should back the VLM. Build (and cache) the native ``transformers`` Florence-2 config that backs the VLM.
``florence_config`` may be given either in the native ``transformers`` format or in the
original Microsoft remote-code format stored by existing XVLA checkpoints (e.g. with
``dim_embed`` / ``image_pos_embed`` in the vision config); the latter is translated
field-by-field to the native format.
""" """
if self._florence_config_obj is None: if self._florence_config_obj is None:
config_dict = dict(self.florence_config) config_dict = dict(self.florence_config)
if "vision_config" not in config_dict or config_dict["vision_config"] is None: if config_dict.get("vision_config") is None:
raise ValueError("vision_config is required") raise ValueError("vision_config is required")
if config_dict.get("text_config") is None:
if "text_config" not in config_dict or config_dict["text_config"] is None:
raise ValueError("text_config is required") raise ValueError("text_config is required")
self._florence_config_obj = Florence2Config(**config_dict)
vision_config = _translate_vision_config(config_dict["vision_config"])
text_config = dict(config_dict["text_config"])
if text_config.get("model_type", "florence2_language") == "florence2_language":
# The MS remote-code language config is BART, field for field.
text_config["model_type"] = "bart"
kwargs = {
key: config_dict[key]
for key in (
"pad_token_id",
"bos_token_id",
"eos_token_id",
"image_token_id",
"is_encoder_decoder",
"tie_word_embeddings",
)
if key in config_dict
}
self._florence_config_obj = Florence2Config(
vision_config=vision_config, text_config=text_config, **kwargs
)
return self._florence_config_obj return self._florence_config_obj
def validate_features(self) -> None: def validate_features(self) -> None:
File diff suppressed because it is too large Load Diff
+97 -62
View File
@@ -21,18 +21,19 @@ from __future__ import annotations
import builtins import builtins
import logging import logging
import os import os
import re
from collections import deque from collections import deque
from pathlib import Path from pathlib import Path
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
import torch import torch
import torch.nn.functional as F # noqa: N812
from torch import Tensor, nn from torch import Tensor, nn
from lerobot.configs import PreTrainedConfig from lerobot.configs import PreTrainedConfig
from lerobot.utils.constants import ACTION, OBS_LANGUAGE_TOKENS, OBS_STATE from lerobot.utils.constants import ACTION, OBS_LANGUAGE_TOKENS, OBS_STATE
from lerobot.utils.import_utils import _transformers_available, require_package from lerobot.utils.import_utils import _transformers_available, require_package
from ..common.vla_utils import pad_vector, resize_with_pad
from ..pretrained import PreTrainedPolicy, T from ..pretrained import PreTrainedPolicy, T
from ..utils import populate_queues from ..utils import populate_queues
from .action_hub import build_action_space from .action_hub import build_action_space
@@ -41,11 +42,10 @@ from .soft_transformer import SoftPromptedTransformer
# Florence2 config and modeling depend on transformers # Florence2 config and modeling depend on transformers
if TYPE_CHECKING or _transformers_available: if TYPE_CHECKING or _transformers_available:
from .configuration_florence2 import Florence2Config from transformers import Florence2Config, Florence2Model
from .modeling_florence2 import Florence2ForConditionalGeneration
else: else:
Florence2Config = None Florence2Config = None
Florence2ForConditionalGeneration = None Florence2Model = None
class XVLAModel(nn.Module): class XVLAModel(nn.Module):
@@ -83,15 +83,11 @@ class XVLAModel(nn.Module):
self.dim_action = self.action_space.dim_action self.dim_action = self.action_space.dim_action
self.dim_proprio = proprio_dim self.dim_proprio = proprio_dim
self.vlm = Florence2ForConditionalGeneration(florence_config) self.vlm = Florence2Model(florence_config)
if hasattr(self.vlm, "language_model"): # XVLA only uses the encoder-side path of Florence-2; drop the text decoder entirely.
lm = self.vlm.language_model del self.vlm.language_model.decoder
if hasattr(lm, "model") and hasattr(lm.model, "decoder"):
del lm.model.decoder
if hasattr(lm, "lm_head"):
del lm.lm_head
projection_dim = getattr(self.vlm.config, "projection_dim", None) projection_dim = getattr(florence_config.vision_config, "projection_dim", None)
if projection_dim is None: if projection_dim is None:
raise ValueError("Florence2 config must provide `projection_dim` for multimodal fusion.") raise ValueError("Florence2 config must provide `projection_dim` for multimodal fusion.")
@@ -143,12 +139,12 @@ class XVLAModel(nn.Module):
if self.config.freeze_language_encoder and hasattr(self.vlm, "language_model"): if self.config.freeze_language_encoder and hasattr(self.vlm, "language_model"):
lm = self.vlm.language_model lm = self.vlm.language_model
# Freeze encoder # Freeze encoder
if hasattr(lm, "model") and hasattr(lm.model, "encoder"): if hasattr(lm, "encoder"):
for param in lm.model.encoder.parameters(): for param in lm.encoder.parameters():
param.requires_grad = False param.requires_grad = False
# Freeze shared embeddings # Freeze shared embeddings
if hasattr(lm, "model") and hasattr(lm.model, "shared"): if hasattr(lm, "shared"):
for param in lm.model.shared.parameters(): for param in lm.shared.parameters():
param.requires_grad = False param.requires_grad = False
# Freeze or unfreeze policy transformer # Freeze or unfreeze policy transformer
@@ -179,19 +175,19 @@ class XVLAModel(nn.Module):
raise ValueError("At least one image view must be valid per batch.") raise ValueError("At least one image view must be valid per batch.")
valid_images = flat_images[flat_mask] valid_images = flat_images[flat_mask]
valid_feats = self.vlm._encode_image(valid_images) valid_feats = self.vlm.get_image_features(valid_images).pooler_output
tokens_per_view, hidden_dim = valid_feats.shape[1:] tokens_per_view, hidden_dim = valid_feats.shape[1:]
image_features = valid_feats.new_zeros((batch_size * num_views, tokens_per_view, hidden_dim)) image_features = valid_feats.new_zeros((batch_size * num_views, tokens_per_view, hidden_dim))
image_features[flat_mask] = valid_feats image_features[flat_mask] = valid_feats
image_features = image_features.view(batch_size, num_views, tokens_per_view, hidden_dim) image_features = image_features.view(batch_size, num_views, tokens_per_view, hidden_dim)
inputs_embeds = self.vlm.get_input_embeddings()(input_ids) inputs_embeds = self.vlm.get_input_embeddings()(input_ids)
merged_embeds, attention_mask = self.vlm._merge_input_ids_with_image_features(
image_features[:, 0],
inputs_embeds,
)
enc_out = self.vlm.language_model.model.encoder( # XVLA prepends the primary view's image tokens to the text embeddings and attends to everything.
merged_embeds = torch.cat([image_features[:, 0], inputs_embeds], dim=1)
attention_mask = torch.ones(merged_embeds.shape[:2], dtype=torch.long, device=merged_embeds.device)
enc_out = self.vlm.language_model.encoder(
attention_mask=attention_mask, attention_mask=attention_mask,
inputs_embeds=merged_embeds, inputs_embeds=merged_embeds,
)[0] )[0]
@@ -310,7 +306,7 @@ class XVLAPolicy(PreTrainedPolicy):
state = batch[OBS_STATE] state = batch[OBS_STATE]
if state.ndim > 2: if state.ndim > 2:
state = state[:, -1, :] state = state[:, -1, :]
return pad_vector(state, self.model.dim_proprio) return pad_vector(state, self.model.dim_proprio, truncate=True)
def _prepare_images(self, batch: dict[str, Tensor]) -> tuple[Tensor, Tensor]: def _prepare_images(self, batch: dict[str, Tensor]) -> tuple[Tensor, Tensor]:
present_img_keys = [key for key in self.config.image_features if key in batch] present_img_keys = [key for key in self.config.image_features if key in batch]
@@ -325,7 +321,7 @@ class XVLAPolicy(PreTrainedPolicy):
for key in present_img_keys: for key in present_img_keys:
img = batch[key][:, -1] if batch[key].ndim == 5 else batch[key] img = batch[key][:, -1] if batch[key].ndim == 5 else batch[key]
if self.config.resize_imgs_with_padding is not None: if self.config.resize_imgs_with_padding is not None:
img = resize_with_pad(img, *self.config.resize_imgs_with_padding) img = resize_with_pad(img, *self.config.resize_imgs_with_padding, pad_value=0.0)
images.append(img) images.append(img)
masks.append(torch.ones(img.size(0), dtype=torch.bool, device=img.device)) masks.append(torch.ones(img.size(0), dtype=torch.bool, device=img.device))
@@ -375,7 +371,7 @@ class XVLAPolicy(PreTrainedPolicy):
actions = actions.unsqueeze(1) actions = actions.unsqueeze(1)
actions = pad_tensor_along_dim(actions, self.config.chunk_size, dim=1) actions = pad_tensor_along_dim(actions, self.config.chunk_size, dim=1)
if actions.shape[-1] != self.model.dim_action: if actions.shape[-1] != self.model.dim_action:
actions = pad_vector(actions, self.model.dim_action) actions = pad_vector(actions, self.model.dim_action, truncate=True)
return actions return actions
def _build_model_inputs(self, batch: dict[str, Tensor]) -> dict[str, Tensor]: def _build_model_inputs(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
@@ -488,13 +484,24 @@ class XVLAPolicy(PreTrainedPolicy):
raise FileNotFoundError(f"model.safetensors not found on the Hub at {model_id}") from e raise FileNotFoundError(f"model.safetensors not found on the Hub at {model_id}") from e
logging.info(f"Loading checkpoint from {model_file}") logging.info(f"Loading checkpoint from {model_file}")
# step 3: load state dict # step 3: load state dict, remapping checkpoints saved with the old vendored
# Florence-2 module layout to the native transformers layout
# (see openpi model.py `_fix_pytorch_state_dict_keys` / pi0 for the same pattern)
state_dict = safetensors.torch.load_file(model_file) state_dict = safetensors.torch.load_file(model_file)
encoder_key = "model.vlm.language_model.model.encoder.embed_tokens.weight" if _is_vendored_florence_state_dict(state_dict):
shared_key = "model.vlm.language_model.model.shared.weight" logging.info(
if encoder_key in state_dict: "Detected XVLA checkpoint with the old vendored Florence-2 layout; "
state_dict[shared_key] = state_dict[encoder_key] "remapping keys to the native transformers layout."
# or deepcopy )
state_dict = _remap_vendored_florence_state_dict(state_dict)
# safetensors deduplicates tied tensors on save: restore whichever alias of the
# shared/encoder token embedding is missing
shared_key = "model.vlm.language_model.shared.weight"
embed_key = "model.vlm.language_model.encoder.embed_tokens.weight"
if shared_key in state_dict and embed_key not in state_dict:
state_dict[embed_key] = state_dict[shared_key]
elif embed_key in state_dict and shared_key not in state_dict:
state_dict[shared_key] = state_dict[embed_key]
# step 4: load into instance # step 4: load into instance
instance.load_state_dict(state_dict, strict=True) instance.load_state_dict(state_dict, strict=True)
logging.info("Loaded XVLA checkpoint") logging.info("Loaded XVLA checkpoint")
@@ -506,41 +513,69 @@ class XVLAPolicy(PreTrainedPolicy):
return instance return instance
def resize_with_pad(img: torch.Tensor, height: int, width: int, pad_value: float = 0.0) -> torch.Tensor: def _is_vendored_florence_state_dict(state_dict: dict[str, Tensor], prefix: str = "model.vlm.") -> bool:
if img.ndim != 4: """Detect XVLA checkpoints saved with the old vendored (Microsoft remote-code) Florence-2
raise ValueError(f"(b,c,h,w) expected, but got {img.shape}") module layout by their signature keys."""
return f"{prefix}image_projection" in state_dict or any(
current_height, current_width = img.shape[2:] key.startswith(f"{prefix}language_model.model.") for key in state_dict
if current_height == height and current_width == width:
return img
ratio = max(current_width / width, current_height / height)
resized_height = int(current_height / ratio)
resized_width = int(current_width / ratio)
resized_img = F.interpolate(
img, size=(resized_height, resized_width), mode="bilinear", align_corners=False
) )
pad_height = max(0, height - resized_height)
pad_width = max(0, width - resized_width)
padded_img = F.pad(resized_img, (pad_width, 0, pad_height, 0), value=pad_value)
return padded_img
def _remap_vendored_florence_state_dict(
state_dict: dict[str, Tensor], prefix: str = "model.vlm."
) -> dict[str, Tensor]:
"""Remap a state dict from the vendored (Microsoft remote-code) Florence-2 layout to the
native ``transformers.models.florence2`` layout.
def pad_vector(vector: Tensor, new_dim: int) -> Tensor: Only keys under ``prefix`` are rewritten; everything else passes through unchanged.
if vector.shape[-1] == new_dim: """
return vector vision = re.escape(prefix) + r"vision_tower\."
if new_dim == 0: block = vision + r"blocks\.(\d+)\.(\d+)\.(spatial_block|channel_block)\."
shape = list(vector.shape) new_block = prefix + r"vision_tower.blocks.\1.\2.\3."
shape[-1] = 0 rules: list[tuple[str, str]] = [
return vector.new_zeros(*shape) # DaViT stem: ConvEmbed.proj -> Florence2VisionConvEmbed.conv
shape = list(vector.shape) (vision + r"convs\.(\d+)\.proj\.", prefix + r"vision_tower.convs.\1.conv."),
current_dim = shape[-1] # DaViT blocks: the PreNorm/Mlp wrappers are flattened in the native implementation
shape[-1] = new_dim (block + r"conv1\.fn\.dw\.", new_block + r"conv1."),
new_vector = vector.new_zeros(*shape) (block + r"conv2\.fn\.dw\.", new_block + r"conv2."),
length = min(current_dim, new_dim) (block + r"(window_attn|channel_attn)\.norm\.", new_block + r"norm1."),
new_vector[..., :length] = vector[..., :length] (block + r"(window_attn|channel_attn)\.fn\.", new_block + r"\4."),
return new_vector (block + r"ffn\.norm\.", new_block + r"norm2."),
(block + r"ffn\.fn\.net\.", new_block + r"ffn."),
# multimodal projection layers moved into a dedicated projector module
(re.escape(prefix) + r"image_proj_norm\.", prefix + r"multi_modal_projector.image_proj_norm."),
(
re.escape(prefix) + r"image_pos_embed\.",
prefix + r"multi_modal_projector.image_position_embed.",
),
(
re.escape(prefix) + r"visual_temporal_embed\.",
prefix + r"multi_modal_projector.visual_temporal_embed.",
),
# language model: Florence2LanguageForConditionalGeneration.model -> BartModel
(re.escape(prefix) + r"language_model\.model\.", prefix + r"language_model."),
]
remapped: dict[str, Tensor] = {}
for key, value in state_dict.items():
if key == f"{prefix}language_model.final_logits_bias":
# generation-only buffer of the vendored language model; the native BartModel has none
continue
if key == f"{prefix}image_projection":
# vendored: nn.Parameter of shape (embed_dim, projection_dim), used as `x @ p`;
# native: nn.Linear(embed_dim, projection_dim, bias=False) whose weight is the transpose
remapped[f"{prefix}multi_modal_projector.image_projection.weight"] = value.transpose(
0, 1
).contiguous()
continue
new_key = key
for pattern, replacement in rules:
new_key, count = re.subn(pattern, replacement, new_key, count=1)
if count:
break
remapped[new_key] = value
return remapped
def pad_tensor_along_dim(tensor: Tensor, target_len: int, dim: int = 1) -> Tensor: def pad_tensor_along_dim(tensor: Tensor, target_len: int, dim: int = 1) -> Tensor: