mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-23 17:56:07 +00:00
refactor(xvla): reuse native Florence2 components (#4089)
This commit is contained in:
@@ -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
@@ -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:
|
||||||
|
|||||||
Reference in New Issue
Block a user