adding cache and cpu device for text encoder to avoid oom

This commit is contained in:
Maxime Ellerbach
2026-07-20 15:24:07 +00:00
parent 5faa956d39
commit c27eb3a9a9
4 changed files with 54 additions and 6 deletions
@@ -202,6 +202,11 @@ class FastWAMConfig(PreTrainedConfig):
base_model_id: str | None = FASTWAM_BASE_MODEL_ID base_model_id: str | None = FASTWAM_BASE_MODEL_ID
tokenizer_max_len: int = 128 tokenizer_max_len: int = 128
load_text_encoder: bool = True load_text_encoder: bool = True
# Device for the frozen ~11GB UMT5-XXL text encoder. `None` keeps it on the main
# policy `device` (default). Set to e.g. "cpu" to keep it off the GPU and save VRAM;
# prompts are then encoded on that device and the resulting embeddings moved to the
# policy device. Trades GPU memory for slower (CPU) text encoding.
text_encoder_device: str | None = None
mot_checkpoint_mixed_attn: bool = False mot_checkpoint_mixed_attn: bool = False
torch_dtype: str = "bfloat16" torch_dtype: str = "bfloat16"
prompt_template: str = ( prompt_template: str = (
@@ -150,6 +150,24 @@ class FastWAMPolicy(PreTrainedPolicy):
def reset(self) -> None: def reset(self) -> None:
self._action_queue: deque[Tensor] = deque([], maxlen=self.config.n_action_steps) self._action_queue: deque[Tensor] = deque([], maxlen=self.config.n_action_steps)
# Per-episode text-embedding cache (mirrors LingBot-VA's `_prompt_embeds`). The task
# is fixed for an episode, so the ~11GB UMT5 encoder runs once on the first chunk and
# the resulting context is reused for every subsequent chunk. Cleared here on reset so
# a new episode's (possibly different) task is re-encoded. Proprio is still appended
# fresh each chunk downstream, so only the text-only context is cached.
self._cached_prompt: Any = None
self._cached_context: Tensor | None = None
self._cached_context_mask: Tensor | None = None
def _encode_prompt_cached(self, prompt: Any) -> tuple[Tensor, Tensor]:
"""Encode `prompt` to `(context, context_mask)`, reusing the cache when the prompt is
unchanged so UMT5 runs at most once per episode (per distinct task)."""
if self._cached_context is None or self._cached_prompt != prompt:
context, context_mask = self.model.encode_prompt(prompt)
self._cached_prompt = prompt
self._cached_context = context
self._cached_context_mask = context_mask
return self._cached_context, self._cached_context_mask
def _batch_to_training_sample(self, batch: dict[str, Tensor]) -> dict[str, Tensor]: def _batch_to_training_sample(self, batch: dict[str, Tensor]) -> dict[str, Tensor]:
"""Adapt a standard LeRobot batch to the FastWAM-native sample that """Adapt a standard LeRobot batch to the FastWAM-native sample that
@@ -222,6 +240,15 @@ class FastWAMPolicy(PreTrainedPolicy):
self.eval() self.eval()
infer_kwargs = _batch_to_infer_kwargs(batch=batch, config=self.config) infer_kwargs = _batch_to_infer_kwargs(batch=batch, config=self.config)
# Encode the task once per episode and reuse it (LingBot-VA parity): swap the raw
# `prompt` for the cached `context`/`context_mask` so `infer_action` skips `encode_prompt`
# and the text encoder isn't re-run every chunk. Skipped when the caller supplies its own
# precomputed `context` (the two are mutually exclusive downstream).
if infer_kwargs.get("context") is None and infer_kwargs.get("prompt") is not None:
context, context_mask = self._encode_prompt_cached(infer_kwargs["prompt"])
infer_kwargs["prompt"] = None
infer_kwargs["context"] = context
infer_kwargs["context_mask"] = context_mask
batch_size = _infer_kwargs_batch_size(infer_kwargs) batch_size = _infer_kwargs_batch_size(infer_kwargs)
if batch_size == 1: if batch_size == 1:
action = _action_from_model_output(self.model.infer_action(**infer_kwargs)) action = _action_from_model_output(self.model.infer_action(**infer_kwargs))
@@ -266,9 +293,10 @@ class FastWAMPolicy(PreTrainedPolicy):
mixtures={"video": video_expert, "action": action_expert}, mixtures={"video": video_expert, "action": action_expert},
mot_checkpoint_mixed_attn=config.mot_checkpoint_mixed_attn, mot_checkpoint_mixed_attn=config.mot_checkpoint_mixed_attn,
) )
text_encoder_device = config.text_encoder_device or device
text_encoder = ( text_encoder = (
load_pretrained_wan_text_encoder( load_pretrained_wan_text_encoder(
model_id=config.text_encoder_model_id, torch_dtype=dtype, device=device model_id=config.text_encoder_model_id, torch_dtype=dtype, device=text_encoder_device
) )
if config.load_text_encoder if config.load_text_encoder
else None else None
@@ -279,6 +307,7 @@ class FastWAMPolicy(PreTrainedPolicy):
mot=mot, mot=mot,
vae=load_pretrained_wan_vae(torch_dtype=dtype, device=device), vae=load_pretrained_wan_vae(torch_dtype=dtype, device=device),
text_encoder=text_encoder, text_encoder=text_encoder,
text_encoder_device=config.text_encoder_device,
tokenizer=build_wan_tokenizer( tokenizer=build_wan_tokenizer(
model_id=config.tokenizer_model_id, tokenizer_max_len=config.tokenizer_max_len model_id=config.tokenizer_model_id, tokenizer_max_len=config.tokenizer_max_len
), ),
+19 -4
View File
@@ -839,6 +839,7 @@ class FastWAM(torch.nn.Module):
text_dim: int | None = None, text_dim: int | None = None,
proprio_dim: int | None = None, proprio_dim: int | None = None,
device: str = "cpu", device: str = "cpu",
text_encoder_device: str | torch.device | None = None,
torch_dtype: torch.dtype = torch.float32, torch_dtype: torch.dtype = torch.float32,
video_train_shift: float = 5.0, video_train_shift: float = 5.0,
video_infer_shift: float = 5.0, video_infer_shift: float = 5.0,
@@ -908,11 +909,21 @@ class FastWAM(torch.nn.Module):
self.infer_scheduler = self.infer_video_scheduler self.infer_scheduler = self.infer_video_scheduler
self.device = torch.device(device) self.device = torch.device(device)
# When pinned (e.g. "cpu"), the frozen text encoder stays on this device instead
# of following the model onto the GPU — `_apply` skips it and `encode_prompt` runs
# it here, moving embeddings back to `self.device`. `None` = follow `self.device`.
self._text_encoder_device = (
torch.device(text_encoder_device) if text_encoder_device is not None else None
)
self.torch_dtype = torch_dtype self.torch_dtype = torch_dtype
self.loss_lambda_video = float(loss_lambda_video) self.loss_lambda_video = float(loss_lambda_video)
self.loss_lambda_action = float(loss_lambda_action) self.loss_lambda_action = float(loss_lambda_action)
self.to(self.device) self.to(self.device)
# `self.to` above (via `_apply`) skips a pinned text encoder; make sure it actually
# sits on the pinned device (it was loaded there, but this is a cheap safety net).
if self.text_encoder is not None and self._text_encoder_device is not None:
self.text_encoder._apply(lambda t: t.to(self._text_encoder_device))
@classmethod @classmethod
def from_wan22_pretrained( def from_wan22_pretrained(
@@ -1003,7 +1014,8 @@ class FastWAM(torch.nn.Module):
# while staying out of `state_dict()` / `parameters()`. # while staying out of `state_dict()` / `parameters()`.
super()._apply(fn, *args, **kwargs) super()._apply(fn, *args, **kwargs)
self.vae._apply(fn) self.vae._apply(fn)
if self.text_encoder is not None: # A pinned text encoder (e.g. on CPU) must NOT follow device moves — leave it put.
if self.text_encoder is not None and self._text_encoder_device is None:
self.text_encoder._apply(fn) self.text_encoder._apply(fn)
return self return self
@@ -1024,9 +1036,12 @@ class FastWAM(torch.nn.Module):
"Prompt encoding requires loaded text encoder/tokenizer. " "Prompt encoding requires loaded text encoder/tokenizer. "
"Set `load_text_encoder=true` or provide precomputed `context/context_mask`." "Set `load_text_encoder=true` or provide precomputed `context/context_mask`."
) )
# Run the encoder on its own device (may be pinned to CPU to save VRAM), then
# move the resulting embeddings/mask to the model device for the DiT.
te_device = self._text_encoder_device or self.device
ids, mask = self.tokenizer(prompt, return_mask=True, add_special_tokens=True) ids, mask = self.tokenizer(prompt, return_mask=True, add_special_tokens=True)
ids = ids.to(self.device) ids = ids.to(te_device)
mask = mask.to(self.device, dtype=torch.bool) mask = mask.to(te_device, dtype=torch.bool)
prompt_emb = self.text_encoder(ids, mask) prompt_emb = self.text_encoder(ids, mask)
seq_lens = mask.gt(0).sum(dim=1).long() seq_lens = mask.gt(0).sum(dim=1).long()
for i, v in enumerate(seq_lens): for i, v in enumerate(seq_lens):
@@ -1034,7 +1049,7 @@ class FastWAM(torch.nn.Module):
# Match FastWAM/Wan2.2 context semantics: padding embeddings are zeroed, # Match FastWAM/Wan2.2 context semantics: padding embeddings are zeroed,
# while cross-attention still sees a fixed-length context. # while cross-attention still sees a fixed-length context.
mask = torch.ones_like(mask) mask = torch.ones_like(mask)
return prompt_emb.to(device=self.device), mask return prompt_emb.to(device=self.device), mask.to(device=self.device)
def _append_proprio_to_context( def _append_proprio_to_context(
self, self,
@@ -573,7 +573,6 @@ class LingBotVAPolicy(PreTrainedPolicy):
if size is None: if size is None:
size = (self.config.height, self.config.width) size = (self.config.height, self.config.width)
img = F.interpolate(img, size=size, mode="bilinear", align_corners=False) img = F.interpolate(img, size=size, mode="bilinear", align_corners=False)
print(f"[lingbot_va] {key} pre-scale range: [{img.min().item():.4f}, {img.max().item():.4f}] (expect ~[0,1])")
img = img * 2.0 - 1.0 img = img * 2.0 - 1.0
return img.unsqueeze(2).to(self.dtype) # [1, C, F=1, H, W] return img.unsqueeze(2).to(self.dtype) # [1, C, F=1, H, W]