Download tensor_batch.py from frisson-labs/Faynt-10M-Base: direct link, hf CLI and curl.
- Browser
- Download file 7 kB
-
https://huggingface.co/frisson-labs/Faynt-10M-Base/resolve/main/tensor_batch.py
- Command line
-
hf download hf://frisson-labs/Faynt-10M-Base/tensor_batch.py
-
curl -L -o tensor_batch.py https://huggingface.co/frisson-labs/Faynt-10M-Base/resolve/main/tensor_batch.py
7 kB
| """Parser-free PyTorch tensor boundary for canonical Melee policy batches. | |
| The nested records mirror ``slippi_ai.types`` at the repository-pinned | |
| slippi-ai revision. Every leaf is a :class:`torch.Tensor`; this module does | |
| not import libmelee, Peppi, slippi-ai, or the earlier replay repository. | |
| Component records use a shared leading shape ``S``. ``S`` is normally | |
| ``[B, T]`` for training and full-sequence inference, and may be ``[B]`` for a | |
| single cached inference step. The one exception is :class:`ItemsBatch`, | |
| whose leaves have shape ``[*S, 15]`` for upstream slots ``item_0`` through | |
| ``item_14``. | |
| The earlier E000 representation is not sufficient to construct this boundary | |
| without enrichment. Preprocessing must retain the processed controller | |
| button mask (not only ``buttons_physical``), follower/Nana state, item slots, | |
| Randall and FoD inputs, and the configured player-name code. It must also | |
| resolve missing values into the exact slippi-ai categorical sentinels before | |
| creating these tensors. Replay parsing and enrichment belong upstream, not in | |
| the model repository. | |
| """ | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from typing import Final | |
| import torch | |
| MAX_ITEMS: Final[int] = 15 | |
| """Number of ordered item slots in the pinned slippi-ai ``Items`` schema.""" | |
| BUTTON_ORDER: Final[tuple[str, ...]] = ("A", "B", "X", "Y", "Z", "L", "R", "D_UP") | |
| """Exact upstream controller-button field order.""" | |
| class ButtonsBatch: | |
| """Named digital controller inputs; every leaf is ``bool [*S]``. | |
| These values must be derived with the pinned slippi-ai controller | |
| semantics from the processed Slippi button mask. E000's physical button | |
| bitset alone is not a lossless substitute. | |
| """ | |
| A: torch.Tensor | |
| B: torch.Tensor | |
| X: torch.Tensor | |
| Y: torch.Tensor | |
| Z: torch.Tensor | |
| L: torch.Tensor | |
| R: torch.Tensor | |
| D_UP: torch.Tensor | |
| class StickBatch: | |
| """One analog stick with ``float32 [*S]`` x and y axes.""" | |
| x: torch.Tensor | |
| y: torch.Tensor | |
| class ControllerBatch: | |
| """Raw slippi-ai controller representation over leading shape ``S``. | |
| ``main_stick`` and ``c_stick`` use the upstream logical coordinate range. | |
| ``shoulder`` is the shared logical shoulder value used by slippi-ai, with | |
| shape ``[*S]`` and floating dtype. The custom_v1 codec consumes this raw | |
| structure and owns all bucketing and decoding. | |
| """ | |
| main_stick: StickBatch | |
| c_stick: StickBatch | |
| shoulder: torch.Tensor | |
| buttons: ButtonsBatch | |
| class NanaBatch: | |
| """Ice Climbers follower state, matching upstream field order. | |
| All leaves have shape ``[*S]``. Boolean fields use ``torch.bool``; | |
| categorical/count fields use an integer dtype; positions and shield use a | |
| floating dtype. When Nana is absent, ``exists`` is false and all other | |
| leaves still contain the exact upstream missing/default representation. | |
| """ | |
| exists: torch.Tensor | |
| percent: torch.Tensor | |
| facing: torch.Tensor | |
| x: torch.Tensor | |
| y: torch.Tensor | |
| action: torch.Tensor | |
| invulnerable: torch.Tensor | |
| character: torch.Tensor | |
| jumps_left: torch.Tensor | |
| shield_strength: torch.Tensor | |
| on_ground: torch.Tensor | |
| class PlayerBatch: | |
| """One player slot, matching pinned ``slippi_ai.types.Player``. | |
| Scalar leaves have shape ``[*S]``. ``controller`` is retained for exact | |
| game-schema parity even though the policy's controlled previous/current | |
| action is supplied separately as :attr:`PolicyBatch.controller_t`. | |
| """ | |
| percent: torch.Tensor | |
| facing: torch.Tensor | |
| x: torch.Tensor | |
| y: torch.Tensor | |
| action: torch.Tensor | |
| invulnerable: torch.Tensor | |
| character: torch.Tensor | |
| jumps_left: torch.Tensor | |
| shield_strength: torch.Tensor | |
| on_ground: torch.Tensor | |
| controller: ControllerBatch | |
| nana: NanaBatch | |
| class RandallBatch: | |
| """Yoshi's Story Randall position, ``float32 [*S]`` per coordinate.""" | |
| x: torch.Tensor | |
| y: torch.Tensor | |
| class FoDPlatformsBatch: | |
| """Fountain of Dreams platform heights, ``float32 [*S]``.""" | |
| left: torch.Tensor | |
| right: torch.Tensor | |
| class ItemsBatch: | |
| """The 15 ordered upstream item slots in a stacked tensor layout. | |
| Each leaf has shape ``[*S, MAX_ITEMS]``. The final dimension maps directly | |
| to ``item_0`` through ``item_14``. ``exists`` uses ``torch.bool``; | |
| ``type`` and ``state`` use integer dtypes; x and y use floating dtypes. | |
| Values in unused slots must match the upstream defaults and are ignored | |
| when ``exists`` is false. | |
| """ | |
| exists: torch.Tensor | |
| type: torch.Tensor | |
| state: torch.Tensor | |
| x: torch.Tensor | |
| y: torch.Tensor | |
| class GameStateBatch: | |
| """Full slippi-ai game state over leading shape ``S``. | |
| Perspective preprocessing makes ``p0`` the controlled/self player and | |
| ``p1`` the opponent. ``stage`` is an integer categorical tensor with | |
| shape ``[*S]``. The remaining fields preserve the pinned upstream nesting. | |
| """ | |
| p0: PlayerBatch | |
| p1: PlayerBatch | |
| stage: torch.Tensor | |
| randall: RandallBatch | |
| fod_platforms: FoDPlatformsBatch | |
| items: ItemsBatch | |
| class PolicyBatch: | |
| """Aligned behavior-cloning batch consumed by the policy and its loss. | |
| Every temporal leaf has batch-major shape ``[B, T, ...]``. | |
| ``game_state_t`` and ``controller_t`` are inputs at frame ``t``. | |
| ``controller_t_plus_1`` is already aligned by preprocessing and is the | |
| target for that same tensor position. Model code must not shift it again. | |
| ``reset_mask[b, t]`` means a new replay segment begins before position t. | |
| ``padding_mask[b, t]`` is true for a real, attention-valid frame. | |
| ``valid_position_mask[b, t]`` is true only when the aligned next-frame | |
| controller target is valid for loss. It must be false at replay ends, | |
| across raw-frame gaps, and wherever required source or target fields are | |
| missing. All three masks use ``torch.bool [B, T]``. | |
| ``player_name`` is the optional upstream categorical name code with shape | |
| ``[B, T]``. It may be ``None`` when player-name conditioning is disabled. | |
| """ | |
| game_state_t: GameStateBatch | |
| controller_t: ControllerBatch | |
| controller_t_plus_1: ControllerBatch | |
| reset_mask: torch.Tensor | |
| padding_mask: torch.Tensor | |
| valid_position_mask: torch.Tensor | |
| player_name: torch.Tensor | None = None | |
| __all__ = [ | |
| "BUTTON_ORDER", | |
| "MAX_ITEMS", | |
| "ButtonsBatch", | |
| "ControllerBatch", | |
| "FoDPlatformsBatch", | |
| "GameStateBatch", | |
| "ItemsBatch", | |
| "NanaBatch", | |
| "PlayerBatch", | |
| "PolicyBatch", | |
| "RandallBatch", | |
| "StickBatch", | |
| ] | |