| """Datasets for Vietnamese Sign Language (Full_TriVis) T2M-GPT training. |
| |
| Backed by the memmaps written by `prepare_vsl_data.py`. Three dataset classes, |
| mirroring the roles of dataset_VQ / dataset_tokenize / dataset_TM_train in the |
| original T2M-GPT: |
| |
| VSLVQDataset random fixed-length windows -> stage-1 (VQ-VAE) training |
| VSLTokenizeDataset whole clips -> encoding clips to token ids |
| VSLText2TokenDataset (text, token ids) -> stage-2 (GPT) training |
| |
| Motion vector = 128 DWPose keypoints x (x,y) = 256 dims, frame-normalized |
| [0,1], z-normalized with the train-split mean/std. A per-keypoint validity mask |
| travels with every sample so the losses can ignore undetected keypoints. |
| """ |
| import json |
| import os |
| import random |
|
|
| import numpy as np |
| import torch |
| from torch.utils import data |
|
|
| from dataset.layout import Layout |
|
|
| |
| NKP = 128 |
| DIM = NKP * 2 |
| BODY, FACE, LH, RH = (0, 18), (18, 86), (86, 107), (107, 128) |
|
|
|
|
| HAND_NKP = 21 |
| FINGERTIPS = (4, 8, 12, 16, 20) |
|
|
|
|
| def kp_weights(body=1.0, face=0.5, hand=3.0, layout=None, finger=None, fingertip=None): |
| """Per-dim reconstruction weights: hands matter most for sign language. |
| |
| Each hand block is 21 keypoints in canonical order -- local 0 is the wrist |
| (verified: it coincides with the body wrist to 0.015 frame-widths), then |
| thumb/index/middle/ring/pinky x 4 joints, so local 4/8/12/16/20 are the tips. |
| |
| `finger` overrides `hand` on local 1..20 (every joint but the wrist) and |
| `fingertip` overrides both on the five tips. Both default to None, which |
| reproduces the original uniform-per-hand weighting exactly. |
| """ |
| layout = layout or Layout.full() |
| g = layout.groups |
| w = np.ones(layout.n_kpts, np.float32) |
| w[g['body'][0]:g['body'][1]] = body |
| w[g['face'][0]:g['face'][1]] = face |
| for side in ('lhand', 'rhand'): |
| a, b = g[side] |
| w[a:b] = hand |
| if finger is None and fingertip is None: |
| continue |
| |
| |
| if b - a != HAND_NKP: |
| raise ValueError( |
| f"--w-finger/--w-fingertip need a {HAND_NKP}-keypoint {side}, " |
| f"but layout {layout.name!r} has {b - a}") |
| if finger is not None: |
| w[a + 1:b] = finger |
| if fingertip is not None: |
| for t in FINGERTIPS: |
| w[a + t] = fingertip |
| return np.repeat(w, 2) |
|
|
|
|
| class VSLStore: |
| """Memmap-backed access to one split's frames + per-clip index.""" |
|
|
| def __init__(self, data_dir, split): |
| self.data_dir = data_dir |
| self.split = split |
| self.layout = Layout.load(data_dir) |
| self.n_kpts = self.layout.n_kpts |
| self.dim = self.layout.dim |
| self.xy = np.load(os.path.join(data_dir, f"{split}_xy.npy"), mmap_mode="r") |
| self.valid = np.load(os.path.join(data_dir, f"{split}_valid.npy"), mmap_mode="r") |
| with open(os.path.join(data_dir, f"{split}_index.json"), encoding="utf-8") as f: |
| self.index = json.load(f) |
| self.mean = np.load(os.path.join(data_dir, "mean.npy")) |
| self.std = np.load(os.path.join(data_dir, "std.npy")) |
|
|
| def __len__(self): |
| return len(self.index) |
|
|
| def get(self, i, start=0, length=None): |
| """Return (motion [L,256] float32 z-normalized, mask [L,256] float32).""" |
| c = self.index[i] |
| length = c["length"] if length is None else length |
| s = c["start"] + start |
| xy = np.asarray(self.xy[s:s + length], dtype=np.float32) |
| vd = np.asarray(self.valid[s:s + length], dtype=np.float32) |
| motion = (xy - self.mean) / self.std |
| mask = np.repeat(vd, 2, axis=1) |
| return motion, mask |
|
|
| def inv_transform(self, motion): |
| """z-normalized -> raw frame-normalized [0,1] coordinates.""" |
| return motion * self.std + self.mean |
|
|
|
|
| class VSLVQDataset(data.Dataset): |
| """Random `window_size`-frame windows, one per clip per epoch.""" |
|
|
| def __init__(self, data_dir, split="train", window_size=64): |
| self.store = VSLStore(data_dir, split) |
| self.window_size = window_size |
| self.items = [i for i, c in enumerate(self.store.index) |
| if c["length"] >= window_size] |
| print(f"[VSLVQDataset:{split}] {len(self.items)}/{len(self.store.index)} clips " |
| f"with >= {window_size} frames") |
|
|
| def __len__(self): |
| return len(self.items) |
|
|
| def __getitem__(self, k): |
| i = self.items[k] |
| T = self.store.index[i]["length"] |
| start = random.randint(0, T - self.window_size) |
| motion, mask = self.store.get(i, start, self.window_size) |
| return torch.from_numpy(motion), torch.from_numpy(mask) |
|
|
|
|
| class VSLFixedWindowDataset(data.Dataset): |
| """Deterministic windows (stride) — used for val reconstruction so the |
| reported number does not move with the random seed.""" |
|
|
| def __init__(self, data_dir, split="val", window_size=64, stride=64, max_windows=0): |
| self.store = VSLStore(data_dir, split) |
| self.window_size = window_size |
| self.items = [] |
| for i, c in enumerate(self.store.index): |
| T = c["length"] |
| if T < window_size: |
| continue |
| for s in range(0, T - window_size + 1, stride): |
| self.items.append((i, s)) |
| if max_windows and len(self.items) > max_windows: |
| rng = random.Random(0) |
| self.items = rng.sample(self.items, max_windows) |
| print(f"[VSLFixedWindowDataset:{split}] {len(self.items)} windows") |
|
|
| def __len__(self): |
| return len(self.items) |
|
|
| def __getitem__(self, k): |
| i, s = self.items[k] |
| motion, mask = self.store.get(i, s, self.window_size) |
| return torch.from_numpy(motion), torch.from_numpy(mask) |
|
|
|
|
| class VSLTokenizeDataset(data.Dataset): |
| """Whole clips, length trimmed to a multiple of `unit_length`, for turning |
| the corpus into VQ token sequences (batch_size must be 1: variable length).""" |
|
|
| def __init__(self, data_dir, split, unit_length=4, max_frames=0): |
| self.store = VSLStore(data_dir, split) |
| self.unit_length = unit_length |
| self.max_frames = max_frames |
| self.items = [i for i, c in enumerate(self.store.index) |
| if c["length"] >= unit_length] |
| print(f"[VSLTokenizeDataset:{split}] {len(self.items)} clips") |
|
|
| def __len__(self): |
| return len(self.items) |
|
|
| def __getitem__(self, k): |
| i = self.items[k] |
| c = self.store.index[i] |
| T = c["length"] |
| if self.max_frames: |
| T = min(T, self.max_frames) |
| T = (T // self.unit_length) * self.unit_length |
| motion, mask = self.store.get(i, 0, T) |
| return torch.from_numpy(motion), c["name"], k |
|
|
|
|
| class VSLText2TokenDataset(data.Dataset): |
| """(text, VQ token sequence) pairs for stage-2 GPT training. |
| |
| Token files are the .npy produced by tokenize_vsl.py, one per clip. |
| Sequences are terminated with `end_idx` and padded with `pad_idx`, exactly |
| as T2M-GPT's Text2MotionDataset does. |
| """ |
|
|
| def __init__(self, data_dir, token_dir, split, codebook_size, |
| max_tokens=128, text_field="gloss", augment_crop=True): |
| self.store = VSLStore(data_dir, split) |
| self.token_dir = token_dir |
| self.end_idx = codebook_size |
| self.pad_idx = codebook_size + 1 |
| self.max_tokens = max_tokens |
| self.text_field = text_field |
| self.augment_crop = augment_crop |
|
|
| self.items = [] |
| missing = 0 |
| for i, c in enumerate(self.store.index): |
| p = os.path.join(token_dir, c["name"] + ".npy") |
| if not os.path.exists(p): |
| missing += 1 |
| continue |
| self.items.append((i, p)) |
| print(f"[VSLText2TokenDataset:{split}] {len(self.items)} pairs " |
| f"({missing} missing token files)") |
|
|
| def __len__(self): |
| return len(self.items) |
|
|
| def __getitem__(self, k): |
| i, p = self.items[k] |
| c = self.store.index[i] |
| tokens = np.load(p).reshape(-1).astype(np.int64) |
|
|
| |
| if self.augment_crop and len(tokens) > 2 and np.random.rand() < 1.0 / 3: |
| if np.random.rand() < 0.5: |
| tokens = tokens[:-1] |
| else: |
| tokens = tokens[1:] |
|
|
| if len(tokens) > self.max_tokens - 1: |
| tokens = tokens[: self.max_tokens - 1] |
| n = len(tokens) |
| out = np.full(self.max_tokens, self.pad_idx, dtype=np.int64) |
| out[:n] = tokens |
| out[n] = self.end_idx |
| return c[self.text_field], out, n |
|
|
|
|
| def cycle(iterable): |
| while True: |
| for x in iterable: |
| yield x |
|
|