ZibinDong's picture
Upload pretrained ActionCodec2 artifact
fee0e43 verified
Raw History Blame Contribute Delete
5.76 kB
"""Explicit action-token constraints for Hugging Face generate()."""
from __future__ import annotations
import torch
from transformers import LogitsProcessor
from ..routing.cache import BoundedCache
class ActionCodec2LogitsProcessor(LogitsProcessor):
"""Constrain generated action IDs, then allow EOS only at complete coverage.
Args:
codec: Fitted ActionCodec2 instance; its current grammar is captured.
horizon: Number of source action frames requested.
fps: Explicit source/output sampling rate in Hz.
prompt_length: Number of leading input IDs to ignore, including padding.
For encoder-decoder models this is the decoder prompt length.
eos_token_id: Model EOS ID, outside the action vocabulary interval.
token_offset: Explicit contiguous offset of codec IDs in model vocabulary.
No vocabulary mapping is inferred or installed into the model.
pad_token_id: Optional model padding ID used after a completed EOS.
Each call takes input_ids (B,L) and scores (B,V), including expanded beam rows.
Prefixes move to CPU once per call; device masks are shared by grammar state
in a bounded cache. The original scores tensor is never modified.
"""
def __init__(
self,
codec,
horizon,
*,
fps,
prompt_length,
eos_token_id,
token_offset=0,
pad_token_id=None,
):
for name, value in {
"prompt_length": prompt_length,
"eos_token_id": eos_token_id,
"token_offset": token_offset,
}.items():
if type(value) is not int or value < 0:
raise ValueError(f"{name} must be a nonnegative integer")
if pad_token_id is not None and (
type(pad_token_id) is not int or pad_token_id < 0
):
raise ValueError("pad_token_id must be a nonnegative integer")
self.grammar = codec.grammar(horizon, fps=fps)
self.prompt_length = prompt_length
self.eos_token_id = eos_token_id
self.pad_token_id = pad_token_id
self.token_offset = token_offset
stop = token_offset + self.grammar.vocab_size
if token_offset <= eos_token_id < stop:
raise ValueError("eos_token_id must be outside the action token interval")
if pad_token_id is not None and token_offset <= pad_token_id < stop:
raise ValueError("pad_token_id must be outside the action token interval")
self._mask_cache = BoundedCache(64)
def _state(self, row):
if self.eos_token_id in row:
end = row.index(self.eos_token_id)
if any(t not in (self.eos_token_id, self.pad_token_id) for t in row[end:]):
raise ValueError(
"only EOS/padding may follow completed action generation"
)
tokens = [t - self.token_offset for t in row[:end]]
if not self.grammar.is_complete(tokens):
raise ValueError("EOS occurred before action coverage was complete")
else:
tokens = [t - self.token_offset for t in row]
state = self.grammar._consume(tokens)
if state is None:
raise ValueError(
"generated prefix violates the action grammar; check prompt_length and token_offset"
)
return state
def _mask(self, state, scores):
key = (state, str(scores.device), scores.shape[-1])
mask = self._mask_cache.get(key)
if mask is None:
mask = torch.zeros(scores.shape[-1], device=scores.device, dtype=torch.bool)
index, local = state
if index == len(self.grammar.segments):
mask[self.eos_token_id] = True
else:
segment = self.grammar.segments[index]
allowed = self.grammar._grammars[index]._next_token_mask_from_state(
local
)
start = self.token_offset + segment.profile.spec.token_offset
mask[start : start + len(allowed)] = torch.as_tensor(
allowed, device=scores.device
)
self._mask_cache[key] = mask
return mask
def __call__(
self, input_ids: torch.LongTensor, scores: torch.FloatTensor
) -> torch.FloatTensor:
if (
input_ids.ndim != 2
or scores.ndim != 2
or input_ids.shape[0] != scores.shape[0]
):
raise ValueError(
"expected input_ids(B,L) and scores(B,V) with the same batch size"
)
if input_ids.shape[1] < self.prompt_length:
raise ValueError("prompt_length exceeds the supplied input length")
required = max(
self.token_offset + self.grammar.vocab_size,
self.eos_token_id + 1,
0 if self.pad_token_id is None else self.pad_token_id + 1,
)
if scores.shape[1] < required:
raise ValueError(
"model vocabulary does not cover configured action/EOS/padding IDs"
)
rows = input_ids[:, self.prompt_length :].detach().cpu().tolist()
output = scores.clone()
groups = {}
for index, row in enumerate(rows):
groups.setdefault(self._state(row), []).append(index)
for state, indices in groups.items():
output[indices] = scores[indices].masked_fill(
~self._mask(state, scores), -torch.inf
)
if not torch.isfinite(output).any(dim=1).all():
raise ValueError(
"all legal action logits were suppressed; check other generation processors"
)
return output