mirror of
https://github.com/huggingface/lerobot.git
synced 2026-07-24 02:06:15 +00:00
refactor(smolvla): reuse shared VLA components (#4064)
* refactor(smolvla): reuse shared VLA components * chore(policies): address review smolvla shared utilities
This commit is contained in:
@@ -61,9 +61,15 @@ import torch.nn.functional as F # noqa: N812
|
|||||||
from torch import Tensor, nn
|
from torch import Tensor, nn
|
||||||
|
|
||||||
from lerobot.utils.constants import ACTION, OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS, OBS_STATE
|
from lerobot.utils.constants import ACTION, OBS_LANGUAGE_ATTENTION_MASK, OBS_LANGUAGE_TOKENS, OBS_STATE
|
||||||
from lerobot.utils.device_utils import get_safe_dtype
|
|
||||||
from lerobot.utils.import_utils import require_package
|
from lerobot.utils.import_utils import require_package
|
||||||
|
|
||||||
|
from ..common.flow_matching import euler_integrate, sample_noise, sample_time_beta
|
||||||
|
from ..common.vla_utils import (
|
||||||
|
create_sinusoidal_pos_embedding,
|
||||||
|
make_att_2d_masks,
|
||||||
|
pad_vector,
|
||||||
|
resize_with_pad,
|
||||||
|
)
|
||||||
from ..pretrained import PreTrainedPolicy
|
from ..pretrained import PreTrainedPolicy
|
||||||
from ..rtc.modeling_rtc import RTCProcessor
|
from ..rtc.modeling_rtc import RTCProcessor
|
||||||
from ..utils import (
|
from ..utils import (
|
||||||
@@ -79,96 +85,6 @@ class ActionSelectKwargs(TypedDict, total=False):
|
|||||||
execution_horizon: int | None
|
execution_horizon: int | None
|
||||||
|
|
||||||
|
|
||||||
def create_sinusoidal_pos_embedding(
|
|
||||||
time: torch.tensor, dimension: int, min_period: float, max_period: float, device="cpu"
|
|
||||||
) -> Tensor:
|
|
||||||
"""Computes sine-cosine positional embedding vectors for scalar positions."""
|
|
||||||
if dimension % 2 != 0:
|
|
||||||
raise ValueError(f"dimension ({dimension}) must be divisible by 2")
|
|
||||||
|
|
||||||
if time.ndim != 1:
|
|
||||||
raise ValueError("The time tensor is expected to be of shape `(batch_size, )`.")
|
|
||||||
|
|
||||||
dtype = get_safe_dtype(torch.float64, device.type)
|
|
||||||
fraction = torch.linspace(0.0, 1.0, dimension // 2, dtype=dtype, device=device)
|
|
||||||
period = min_period * (max_period / min_period) ** fraction
|
|
||||||
|
|
||||||
# Compute the outer product
|
|
||||||
scaling_factor = 1.0 / period * 2 * math.pi
|
|
||||||
sin_input = scaling_factor[None, :] * time[:, None]
|
|
||||||
pos_emb = torch.cat([torch.sin(sin_input), torch.cos(sin_input)], dim=1)
|
|
||||||
return pos_emb
|
|
||||||
|
|
||||||
|
|
||||||
def make_att_2d_masks(pad_masks, att_masks):
|
|
||||||
"""Copied from big_vision.
|
|
||||||
|
|
||||||
Tokens can attend to valid inputs tokens which have a cumulative mask_ar
|
|
||||||
smaller or equal to theirs. This way `mask_ar` int[B, N] can be used to
|
|
||||||
setup several types of attention, for example:
|
|
||||||
|
|
||||||
[[1 1 1 1 1 1]]: pure causal attention.
|
|
||||||
|
|
||||||
[[0 0 0 1 1 1]]: prefix-lm attention. The first 3 tokens can attend between
|
|
||||||
themselves and the last 3 tokens have a causal attention. The first
|
|
||||||
entry could also be a 1 without changing behaviour.
|
|
||||||
|
|
||||||
[[1 0 1 0 1 0 0 1 0 0]]: causal attention between 4 blocks. Tokens of a
|
|
||||||
block can attend all previous blocks and all tokens on the same block.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
input_mask: bool[B, N] true if its part of the input, false if padding.
|
|
||||||
mask_ar: int32[B, N] mask that's 1 where previous tokens cannot depend on
|
|
||||||
it and 0 where it shares the same attention mask as the previous token.
|
|
||||||
"""
|
|
||||||
if att_masks.ndim != 2:
|
|
||||||
raise ValueError(att_masks.ndim)
|
|
||||||
if pad_masks.ndim != 2:
|
|
||||||
raise ValueError(pad_masks.ndim)
|
|
||||||
|
|
||||||
cumsum = torch.cumsum(att_masks, dim=1)
|
|
||||||
att_2d_masks = cumsum[:, None, :] <= cumsum[:, :, None]
|
|
||||||
pad_2d_masks = pad_masks[:, None, :] * pad_masks[:, :, None]
|
|
||||||
att_2d_masks = att_2d_masks & pad_2d_masks
|
|
||||||
return att_2d_masks
|
|
||||||
|
|
||||||
|
|
||||||
def resize_with_pad(img, width, height, pad_value=-1):
|
|
||||||
# assume no-op when width height fits already
|
|
||||||
if img.ndim != 4:
|
|
||||||
raise ValueError(f"(b,c,h,w) expected, but {img.shape}")
|
|
||||||
|
|
||||||
cur_height, cur_width = img.shape[2:]
|
|
||||||
|
|
||||||
ratio = max(cur_width / width, cur_height / height)
|
|
||||||
resized_height = int(cur_height / ratio)
|
|
||||||
resized_width = int(cur_width / ratio)
|
|
||||||
resized_img = F.interpolate(
|
|
||||||
img, size=(resized_height, resized_width), mode="bilinear", align_corners=False
|
|
||||||
)
|
|
||||||
|
|
||||||
pad_height = max(0, int(height - resized_height))
|
|
||||||
pad_width = max(0, int(width - resized_width))
|
|
||||||
|
|
||||||
# pad on left and top of image
|
|
||||||
padded_img = F.pad(resized_img, (pad_width, 0, pad_height, 0), value=pad_value)
|
|
||||||
return padded_img
|
|
||||||
|
|
||||||
|
|
||||||
def pad_vector(vector, new_dim):
|
|
||||||
"""Can be (batch_size x sequence_length x features_dimension)
|
|
||||||
or (batch_size x features_dimension)
|
|
||||||
"""
|
|
||||||
if vector.shape[-1] == new_dim:
|
|
||||||
return vector
|
|
||||||
shape = list(vector.shape)
|
|
||||||
current_dim = shape[-1]
|
|
||||||
shape[-1] = new_dim
|
|
||||||
new_vector = torch.zeros(*shape, dtype=vector.dtype, device=vector.device)
|
|
||||||
new_vector[..., :current_dim] = vector
|
|
||||||
return new_vector
|
|
||||||
|
|
||||||
|
|
||||||
def normalize(x, min_val, max_val):
|
def normalize(x, min_val, max_val):
|
||||||
return (x - min_val) / (max_val - min_val)
|
return (x - min_val) / (max_val - min_val)
|
||||||
|
|
||||||
@@ -429,7 +345,13 @@ class SmolVLAPolicy(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, pad_value=0)
|
# SmolVLA stores the target as (width, height); the shared helper expects (height, width).
|
||||||
|
img = resize_with_pad(
|
||||||
|
img,
|
||||||
|
self.config.resize_imgs_with_padding[1],
|
||||||
|
self.config.resize_imgs_with_padding[0],
|
||||||
|
pad_value=0,
|
||||||
|
)
|
||||||
|
|
||||||
# Normalize from range [0,1] to [-1,1] as expacted by siglip
|
# Normalize from range [0,1] to [-1,1] as expacted by siglip
|
||||||
img = img * 2.0 - 1.0
|
img = img * 2.0 - 1.0
|
||||||
@@ -619,20 +541,10 @@ class VLAFlowMatching(nn.Module):
|
|||||||
params.requires_grad = self.config.train_state_proj
|
params.requires_grad = self.config.train_state_proj
|
||||||
|
|
||||||
def sample_noise(self, shape, device):
|
def sample_noise(self, shape, device):
|
||||||
noise = torch.normal(
|
return sample_noise(shape, device)
|
||||||
mean=0.0,
|
|
||||||
std=1.0,
|
|
||||||
size=shape,
|
|
||||||
dtype=torch.float32,
|
|
||||||
device=device,
|
|
||||||
)
|
|
||||||
return noise
|
|
||||||
|
|
||||||
def sample_time(self, bsize, device):
|
def sample_time(self, bsize, device):
|
||||||
beta_dist = torch.distributions.Beta(concentration1=1.5, concentration0=1.0)
|
return sample_time_beta(bsize, device, alpha=1.5, beta=1.0, scale=0.999, offset=0.001)
|
||||||
time_beta = beta_dist.sample((bsize,)).to(device=device, dtype=torch.float32)
|
|
||||||
time = time_beta * 0.999 + 0.001
|
|
||||||
return time
|
|
||||||
|
|
||||||
def embed_prefix(
|
def embed_prefix(
|
||||||
self, images, img_masks, lang_tokens, lang_masks, state: torch.Tensor = None
|
self, images, img_masks, lang_tokens, lang_masks, state: torch.Tensor = None
|
||||||
@@ -800,7 +712,6 @@ class VLAFlowMatching(nn.Module):
|
|||||||
past_key_values=None,
|
past_key_values=None,
|
||||||
inputs_embeds=[prefix_embs, suffix_embs],
|
inputs_embeds=[prefix_embs, suffix_embs],
|
||||||
use_cache=False,
|
use_cache=False,
|
||||||
fill_kv_cache=False,
|
|
||||||
)
|
)
|
||||||
suffix_out = suffix_out[:, -self.config.chunk_size :]
|
suffix_out = suffix_out[:, -self.config.chunk_size :]
|
||||||
# Original openpi code, upcast attention output
|
# Original openpi code, upcast attention output
|
||||||
@@ -839,46 +750,24 @@ class VLAFlowMatching(nn.Module):
|
|||||||
past_key_values=None,
|
past_key_values=None,
|
||||||
inputs_embeds=[prefix_embs, None],
|
inputs_embeds=[prefix_embs, None],
|
||||||
use_cache=self.config.use_cache,
|
use_cache=self.config.use_cache,
|
||||||
fill_kv_cache=True,
|
|
||||||
)
|
)
|
||||||
num_steps = self.config.num_steps
|
num_steps = self.config.num_steps
|
||||||
dt = -1.0 / num_steps
|
|
||||||
|
|
||||||
x_t = noise
|
return euler_integrate(
|
||||||
for step in range(num_steps):
|
lambda input_x_t, current_timestep: self.denoise_step(
|
||||||
time = 1.0 + step * dt
|
x_t=input_x_t,
|
||||||
time_tensor = torch.tensor(time, dtype=torch.float32, device=device).expand(bsize)
|
prefix_pad_masks=prefix_pad_masks,
|
||||||
|
past_key_values=past_key_values,
|
||||||
def denoise_step_partial_call(input_x_t, current_timestep=time_tensor):
|
timestep=current_timestep,
|
||||||
return self.denoise_step(
|
),
|
||||||
x_t=input_x_t,
|
noise,
|
||||||
prefix_pad_masks=prefix_pad_masks,
|
num_steps,
|
||||||
past_key_values=past_key_values,
|
rtc_processor=self.rtc_processor,
|
||||||
timestep=current_timestep,
|
rtc_enabled=self._rtc_enabled(),
|
||||||
)
|
inference_delay=kwargs.get("inference_delay"),
|
||||||
|
prev_chunk_left_over=kwargs.get("prev_chunk_left_over"),
|
||||||
if self._rtc_enabled():
|
execution_horizon=kwargs.get("execution_horizon"),
|
||||||
inference_delay = kwargs.get("inference_delay")
|
)
|
||||||
prev_chunk_left_over = kwargs.get("prev_chunk_left_over")
|
|
||||||
execution_horizon = kwargs.get("execution_horizon")
|
|
||||||
|
|
||||||
v_t = self.rtc_processor.denoise_step(
|
|
||||||
x_t=x_t,
|
|
||||||
prev_chunk_left_over=prev_chunk_left_over,
|
|
||||||
inference_delay=inference_delay,
|
|
||||||
time=time,
|
|
||||||
original_denoise_step_partial=denoise_step_partial_call,
|
|
||||||
execution_horizon=execution_horizon,
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
v_t = denoise_step_partial_call(x_t)
|
|
||||||
|
|
||||||
x_t = x_t + dt * v_t
|
|
||||||
|
|
||||||
if self.rtc_processor is not None and self.rtc_processor.is_debug_enabled():
|
|
||||||
self.rtc_processor.track(time=time, x_t=x_t, v_t=v_t)
|
|
||||||
|
|
||||||
return x_t
|
|
||||||
|
|
||||||
def denoise_step(
|
def denoise_step(
|
||||||
self,
|
self,
|
||||||
@@ -907,8 +796,10 @@ class VLAFlowMatching(nn.Module):
|
|||||||
past_key_values=past_key_values,
|
past_key_values=past_key_values,
|
||||||
inputs_embeds=[None, suffix_embs],
|
inputs_embeds=[None, suffix_embs],
|
||||||
use_cache=self.config.use_cache,
|
use_cache=self.config.use_cache,
|
||||||
fill_kv_cache=False,
|
|
||||||
)
|
)
|
||||||
|
if past_key_values is not None:
|
||||||
|
# Self-attention layers append suffix K/V in place; restore the prefix for the next step.
|
||||||
|
past_key_values.crop(prefix_len)
|
||||||
suffix_out = outputs_embeds[1]
|
suffix_out = outputs_embeds[1]
|
||||||
suffix_out = suffix_out[:, -self.config.chunk_size :]
|
suffix_out = suffix_out[:, -self.config.chunk_size :]
|
||||||
suffix_out = suffix_out.to(dtype=torch.float32)
|
suffix_out = suffix_out.to(dtype=torch.float32)
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ if TYPE_CHECKING or _transformers_available:
|
|||||||
AutoModel,
|
AutoModel,
|
||||||
AutoModelForImageTextToText,
|
AutoModelForImageTextToText,
|
||||||
AutoProcessor,
|
AutoProcessor,
|
||||||
|
DynamicCache,
|
||||||
SmolVLMForConditionalGeneration,
|
SmolVLMForConditionalGeneration,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -33,6 +34,7 @@ else:
|
|||||||
AutoModel = None
|
AutoModel = None
|
||||||
AutoModelForImageTextToText = None
|
AutoModelForImageTextToText = None
|
||||||
AutoProcessor = None
|
AutoProcessor = None
|
||||||
|
DynamicCache = None
|
||||||
SmolVLMForConditionalGeneration = None
|
SmolVLMForConditionalGeneration = None
|
||||||
|
|
||||||
|
|
||||||
@@ -216,9 +218,8 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
batch_size,
|
batch_size,
|
||||||
head_dim,
|
head_dim,
|
||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
fill_kv_cache: bool = True,
|
past_key_values: "DynamicCache | None" = None,
|
||||||
past_key_values=None,
|
) -> "tuple[list[torch.Tensor], DynamicCache | None]":
|
||||||
) -> list[torch.Tensor]:
|
|
||||||
query_states = []
|
query_states = []
|
||||||
key_states = []
|
key_states = []
|
||||||
value_states = []
|
value_states = []
|
||||||
@@ -259,22 +260,16 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
query_states = apply_rope(query_states, position_ids_)
|
query_states = apply_rope(query_states, position_ids_)
|
||||||
key_states = apply_rope(key_states, position_ids_)
|
key_states = apply_rope(key_states, position_ids_)
|
||||||
|
|
||||||
if use_cache and past_key_values is None:
|
|
||||||
past_key_values = {}
|
|
||||||
|
|
||||||
if use_cache:
|
if use_cache:
|
||||||
if fill_kv_cache:
|
# `DynamicCache` stores tensors as [batch, heads, seq, head_dim]; this module works with
|
||||||
past_key_values[layer_idx] = {
|
# [batch, seq, heads, head_dim]. During prefix prefill this stores the (post-RoPE) K/V and
|
||||||
"key_states": key_states,
|
# returns them unchanged; during denoising it appends the suffix K/V and returns
|
||||||
"value_states": value_states,
|
# [prefix; suffix], exactly like the previous hand-rolled dict cache.
|
||||||
}
|
key_states, value_states = past_key_values.update(
|
||||||
else:
|
key_states.transpose(1, 2), value_states.transpose(1, 2), layer_idx
|
||||||
# TODO here, some optimization can be done - similar to a `StaticCache` we can declare the `max_len` before.
|
)
|
||||||
# so we create an empty cache, with just one cuda malloc, and if (in autoregressive case) we reach
|
key_states = key_states.transpose(1, 2)
|
||||||
# the max len, then we (for instance) double the cache size. This implementation already exists
|
value_states = value_states.transpose(1, 2)
|
||||||
# in `transformers`. (molbap)
|
|
||||||
key_states = torch.cat([past_key_values[layer_idx]["key_states"], key_states], dim=1)
|
|
||||||
value_states = torch.cat([past_key_values[layer_idx]["value_states"], value_states], dim=1)
|
|
||||||
|
|
||||||
attention_interface = self.get_attention_interface()
|
attention_interface = self.get_attention_interface()
|
||||||
|
|
||||||
@@ -293,13 +288,12 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
batch_size,
|
batch_size,
|
||||||
head_dim,
|
head_dim,
|
||||||
use_cache: bool = True,
|
use_cache: bool = True,
|
||||||
fill_kv_cache: bool = True,
|
past_key_values: "DynamicCache | None" = None,
|
||||||
past_key_values=None,
|
) -> "tuple[list[torch.Tensor], DynamicCache | None]":
|
||||||
) -> list[torch.Tensor]:
|
|
||||||
attention_interface = self.get_attention_interface()
|
attention_interface = self.get_attention_interface()
|
||||||
|
|
||||||
att_outputs = []
|
att_outputs = []
|
||||||
assert len(inputs_embeds) == 2 or (use_cache and past_key_values is not None and not fill_kv_cache), (
|
assert len(inputs_embeds) == 2 or (use_cache and past_key_values is not None), (
|
||||||
f"Both len(inputs_embeds) == {len(inputs_embeds)} and past_key_values is {past_key_values}"
|
f"Both len(inputs_embeds) == {len(inputs_embeds)} and past_key_values is {past_key_values}"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -332,22 +326,13 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
else:
|
else:
|
||||||
expert_position_id = position_ids
|
expert_position_id = position_ids
|
||||||
|
|
||||||
if use_cache and past_key_values is None:
|
if use_cache and past_key_values is not None:
|
||||||
past_key_values = {}
|
# Cross-attention layers never fill the cache themselves: during the prefix prefill every
|
||||||
|
# layer goes through `forward_attn_layer`, which stores the (post-RoPE) VLM K/V for this
|
||||||
if use_cache:
|
# layer index. Here we only read them back (no concatenation: the expert cross-attends to
|
||||||
if fill_kv_cache:
|
# the fixed prefix). `DynamicCache` stores [batch, heads, seq, head_dim]; transpose back.
|
||||||
past_key_values[layer_idx] = {
|
key_states = past_key_values.layers[layer_idx].keys.transpose(1, 2)
|
||||||
"key_states": key_states,
|
value_states = past_key_values.layers[layer_idx].values.transpose(1, 2)
|
||||||
"value_states": value_states,
|
|
||||||
}
|
|
||||||
else:
|
|
||||||
# TODO here, some optimization can be done - similar to a `StaticCache` we can declare the `max_len` before.
|
|
||||||
# so we create an empty cache, with just one cuda malloc, and if (in autoregressive case) we reach
|
|
||||||
# the max len, then we (for instance) double the cache size. This implementation already exists
|
|
||||||
# in `transformers`. (molbap)
|
|
||||||
key_states = past_key_values[layer_idx]["key_states"]
|
|
||||||
value_states = past_key_values[layer_idx]["value_states"]
|
|
||||||
|
|
||||||
# Expert
|
# Expert
|
||||||
expert_layer = model_layers[1][layer_idx]
|
expert_layer = model_layers[1][layer_idx]
|
||||||
@@ -360,14 +345,15 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
expert_hidden_states = expert_hidden_states.to(dtype=expert_layer.self_attn.q_proj.weight.dtype)
|
expert_hidden_states = expert_hidden_states.to(dtype=expert_layer.self_attn.q_proj.weight.dtype)
|
||||||
expert_query_state = expert_layer.self_attn.q_proj(expert_hidden_states).view(expert_hidden_shape)
|
expert_query_state = expert_layer.self_attn.q_proj(expert_hidden_states).view(expert_hidden_shape)
|
||||||
|
|
||||||
_key_states = key_states.to(dtype=expert_layer.self_attn.k_proj.weight.dtype).view(
|
# reshape (not view): K/V read back from the cache are transposed, hence non-contiguous
|
||||||
|
_key_states = key_states.to(dtype=expert_layer.self_attn.k_proj.weight.dtype).reshape(
|
||||||
*key_states.shape[:2], -1
|
*key_states.shape[:2], -1
|
||||||
)
|
)
|
||||||
expert_key_states = expert_layer.self_attn.k_proj(_key_states).view(
|
expert_key_states = expert_layer.self_attn.k_proj(_key_states).view(
|
||||||
*_key_states.shape[:-1], -1, expert_layer.self_attn.head_dim
|
*_key_states.shape[:-1], -1, expert_layer.self_attn.head_dim
|
||||||
) # k_proj should have same dim as kv
|
) # k_proj should have same dim as kv
|
||||||
|
|
||||||
_value_states = value_states.to(dtype=expert_layer.self_attn.v_proj.weight.dtype).view(
|
_value_states = value_states.to(dtype=expert_layer.self_attn.v_proj.weight.dtype).reshape(
|
||||||
*value_states.shape[:2], -1
|
*value_states.shape[:2], -1
|
||||||
)
|
)
|
||||||
expert_value_states = expert_layer.self_attn.v_proj(_value_states).view(
|
expert_value_states = expert_layer.self_attn.v_proj(_value_states).view(
|
||||||
@@ -416,10 +402,9 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
self,
|
self,
|
||||||
attention_mask: torch.Tensor | None = None,
|
attention_mask: torch.Tensor | None = None,
|
||||||
position_ids: torch.LongTensor | None = None,
|
position_ids: torch.LongTensor | None = None,
|
||||||
past_key_values: list[torch.FloatTensor] | None = None,
|
past_key_values: "DynamicCache | None" = None,
|
||||||
inputs_embeds: list[torch.FloatTensor] = None,
|
inputs_embeds: list[torch.FloatTensor] = None,
|
||||||
use_cache: bool | None = None,
|
use_cache: bool | None = None,
|
||||||
fill_kv_cache: bool | None = None,
|
|
||||||
):
|
):
|
||||||
models = [self.get_vlm_model().text_model, self.lm_expert]
|
models = [self.get_vlm_model().text_model, self.lm_expert]
|
||||||
model_layers = self.get_model_layers(models)
|
model_layers = self.get_model_layers(models)
|
||||||
@@ -431,6 +416,13 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
continue
|
continue
|
||||||
batch_size = hidden_states.shape[0]
|
batch_size = hidden_states.shape[0]
|
||||||
|
|
||||||
|
# Prefix prefill: no cache was passed, so create one and fill it (every layer runs
|
||||||
|
# self-attention over the prefix). When a filled cache is passed (denoising), layers
|
||||||
|
# read from it instead.
|
||||||
|
fill_kv_cache = use_cache and past_key_values is None
|
||||||
|
if fill_kv_cache:
|
||||||
|
past_key_values = DynamicCache()
|
||||||
|
|
||||||
# RMSNorm
|
# RMSNorm
|
||||||
num_layers = self.num_vlm_layers
|
num_layers = self.num_vlm_layers
|
||||||
head_dim = self.vlm.config.text_config.head_dim
|
head_dim = self.vlm.config.text_config.head_dim
|
||||||
@@ -449,7 +441,6 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
batch_size,
|
batch_size,
|
||||||
head_dim,
|
head_dim,
|
||||||
use_cache=use_cache,
|
use_cache=use_cache,
|
||||||
fill_kv_cache=fill_kv_cache,
|
|
||||||
past_key_values=past_key_values,
|
past_key_values=past_key_values,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -462,7 +453,6 @@ class SmolVLMWithExpertModel(nn.Module):
|
|||||||
batch_size,
|
batch_size,
|
||||||
head_dim,
|
head_dim,
|
||||||
use_cache=use_cache,
|
use_cache=use_cache,
|
||||||
fill_kv_cache=fill_kv_cache,
|
|
||||||
past_key_values=past_key_values,
|
past_key_values=past_key_values,
|
||||||
)
|
)
|
||||||
outputs_embeds = []
|
outputs_embeds = []
|
||||||
|
|||||||
Reference in New Issue
Block a user