Instructions to use ZibinDong/ActionCodec2-2nd-order with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use ZibinDong/ActionCodec2-2nd-order with Transformers:
# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("ZibinDong/ActionCodec2-2nd-order", device_map="auto") - Notebooks
- Google Colab
- Kaggle
Download runtime/tokenization/generation.py from ZibinDong/ActionCodec2-2nd-order: direct link, hf CLI and curl.
- Browser
- Download file 5.76 kB
-
https://huggingface.co/ZibinDong/ActionCodec2-2nd-order/resolve/main/runtime/tokenization/generation.py
- Command line
-
hf download hf://ZibinDong/ActionCodec2-2nd-order/runtime/tokenization/generation.py
-
curl -L -o generation.py https://huggingface.co/ZibinDong/ActionCodec2-2nd-order/resolve/main/runtime/tokenization/generation.py
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 | |