File size: 15,937 Bytes
3b18ebd | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 | # -*- coding: utf-8 -*-
"""資料載入、前處理、增強。
原始資料:DATA_ROOT/<編號>.<英文詞>/<classID><sampleID>.npy,每個檔案 shape = (T, 55, 3)。
前處理鏈:NaN 補值 → 以頸點置中、肩寬正規化 → 時間軸重採樣 → 速度/遮罩特徵。
"""
from __future__ import annotations
import glob
import re
from pathlib import Path
import numpy as np
from torch.utils.data import Dataset
import config as C
# ====================================================================== 掃描資料集
def scan_dataset(root: Path):
"""回傳 (檔案路徑清單, 標籤索引清單, 類別名稱清單)。"""
root = Path(root)
if not root.exists():
raise FileNotFoundError(
f"找不到資料夾 {root}\n"
"請先把雲端的 NPY_Dataset 下載到本機,或修改 config.py 的 DATA_ROOT。"
)
def sort_key(name: str):
m = re.match(r"(\d+)", name)
return (int(m.group(1)) if m else 10**9, name)
class_dirs = sorted(
[d for d in root.iterdir() if d.is_dir()], key=lambda d: sort_key(d.name)
)
if not class_dirs:
raise RuntimeError(f"{root} 底下沒有任何類別資料夾")
files, labels, classes = [], [], []
for d in class_dirs:
paths = sorted(glob.glob(str(d / "*.npy")))
if not paths:
print(f"[警告] 類別 {d.name} 沒有 .npy,略過")
continue
classes.append(d.name)
for p in paths:
files.append(p)
labels.append(len(classes) - 1)
return files, np.array(labels, dtype=np.int64), classes
# ====================================================================== 前處理
def _interp_nan_along_time(seq: np.ndarray) -> np.ndarray:
"""對每個關鍵點沿時間軸線性內插補 NaN;整段皆 NaN 的點補 0。"""
T, P, D = seq.shape
flat = seq.reshape(T, P * D).copy()
t = np.arange(T)
for j in range(flat.shape[1]):
col = flat[:, j]
good = ~np.isnan(col)
if good.all():
continue
if good.sum() == 0:
col[:] = 0.0 # 這個點整段沒偵測到(例如全程沒出現的那隻手)
else:
col[~good] = np.interp(t[~good], t[good], col[good])
flat[:, j] = col
return flat.reshape(T, P, D)
def _resample_time(seq: np.ndarray, out_len: int) -> np.ndarray:
"""把 (T, P, D) 沿時間軸線性重採樣成 (out_len, P, D)。"""
T = seq.shape[0]
if T == out_len:
return seq
src = np.linspace(0.0, T - 1, num=T)
dst = np.linspace(0.0, T - 1, num=out_len)
flat = seq.reshape(T, -1)
out = np.empty((out_len, flat.shape[1]), dtype=np.float32)
for j in range(flat.shape[1]):
out[:, j] = np.interp(dst, src, flat[:, j])
return out.reshape(out_len, seq.shape[1], seq.shape[2])
def normalize_sequence(raw: np.ndarray):
"""把單一原始序列標準化成與拍攝距離、身體位置無關的座標。
Args:
raw: (T, 55, C) 原始關鍵點,可能含 NaN。
Returns:
coords: (T, 55, D) 已置中/縮放/補值的座標
mask: (T, 55) 1 = 該格該點原本有偵測到
"""
seq = np.asarray(raw, dtype=np.float32)
if seq.ndim != 3:
raise ValueError(f"預期 (T, P, C) 三維陣列,收到 {seq.shape}")
if C.USE_XY_ONLY and seq.shape[2] >= 2:
seq = seq[:, :, :2]
mask = (~np.isnan(seq).any(axis=2)).astype(np.float32) # (T, P)
# --- 原點:頸點(缺失時退回所有有效點的平均)
if seq.shape[1] > C.CENTER_IDX:
center = seq[:, C.CENTER_IDX, :] # (T, D)
else:
center = np.nanmean(seq, axis=1)
bad = np.isnan(center).any(axis=1)
if bad.any():
fallback = np.nanmean(seq, axis=1)
center = np.where(bad[:, None], fallback, center)
center = np.nan_to_num(center, nan=0.0)
# --- 尺度:肩寬中位數(缺失時退回關鍵點座標的標準差)
if seq.shape[1] > max(C.LSHOULDER_IDX, C.RSHOULDER_IDX):
span = np.linalg.norm(
seq[:, C.LSHOULDER_IDX, :] - seq[:, C.RSHOULDER_IDX, :], axis=-1
)
scale = np.nanmedian(span)
else:
scale = np.nan
if not np.isfinite(scale) or scale < 1e-6:
scale = float(np.nanstd(seq)) or 1.0
seq = (seq - center[:, None, :]) / float(scale)
seq = _interp_nan_along_time(seq)
return seq.astype(np.float32), mask
_PARENTS, _BFS_ORDER = None, None
def _parents():
"""延遲載入骨架樹(避免 data.py 與 graph.py 的循環匯入)。"""
global _PARENTS, _BFS_ORDER
if _PARENTS is None:
from graph import build_parents
_PARENTS, _BFS_ORDER = build_parents(C.NUM_POINTS)
return _PARENTS, _BFS_ORDER
def build_features(coords: np.ndarray, mask: np.ndarray) -> np.ndarray:
"""(T, P, D) + (T, P) → (T, P, Ch) 每個節點自己的特徵向量。
保留「節點」維度,讓 ST-GCN 能直接使用每個骨架節點的特徵。
通道組成(依 config 開關):
座標 x,y 位置資訊,但仍帶有個人習慣
速度 vx,vy 動作方向與快慢
遮罩 mask 這一點這一格是不是補出來的
骨向量 bx,by 指向父節點,與絕對位置無關
關節角度 angle 純角度,與體型、肢長完全無關 ← 跨人泛化的關鍵
"""
T, P, D = coords.shape
parts = [coords]
if C.ADD_VELOCITY:
vel = np.zeros_like(coords)
vel[1:] = coords[1:] - coords[:-1]
parts.append(vel)
if C.ADD_MASK:
parts.append(mask[:, :, None])
if C.ADD_BONE or C.ADD_ANGLE:
parents, _ = _parents()
par = np.where(parents >= 0, parents, np.arange(P))
bone = coords - coords[:, par, :] # (T, P, D),根節點為 0
if C.ADD_BONE:
parts.append(bone)
if C.ADD_ANGLE:
# 該節點的骨向量 與 其父節點的骨向量 之間的夾角餘弦
parent_bone = bone[:, par, :]
n1 = np.linalg.norm(bone, axis=2, keepdims=True)
n2 = np.linalg.norm(parent_bone, axis=2, keepdims=True)
cos = (bone * parent_bone).sum(axis=2, keepdims=True) / np.maximum(n1 * n2, 1e-6)
parts.append(np.clip(cos, -1.0, 1.0))
return np.concatenate(parts, axis=2).astype(np.float32)
def node_channels() -> int:
"""每個節點每一格有幾個通道。"""
d = 2 if C.USE_XY_ONLY else 3
ch = d
if C.ADD_VELOCITY:
ch += d
if C.ADD_MASK:
ch += 1
if C.ADD_BONE:
ch += d
if C.ADD_ANGLE:
ch += 1
return ch
def feature_flags() -> dict:
"""存進 checkpoint,推論時據此還原相同的特徵設定。"""
return dict(USE_XY_ONLY=C.USE_XY_ONLY, ADD_VELOCITY=C.ADD_VELOCITY,
ADD_MASK=C.ADD_MASK, ADD_BONE=C.ADD_BONE, ADD_ANGLE=C.ADD_ANGLE,
SEQ_LEN=C.SEQ_LEN)
def flags_from_ckpt(ckpt: dict) -> dict:
"""從 checkpoint 取出特徵設定,並相容於還沒有骨向量/角度通道的舊檔。"""
if ckpt.get("feature_flags"):
return ckpt["feature_flags"]
return dict(
USE_XY_ONLY=ckpt.get("use_xy_only", True),
ADD_VELOCITY=ckpt.get("add_velocity", True),
ADD_MASK=ckpt.get("add_mask", True),
ADD_BONE=False, ADD_ANGLE=False, # 舊版沒有這兩組通道
SEQ_LEN=ckpt.get("seq_len", C.SEQ_LEN),
)
def apply_feature_flags(ckpt_or_flags):
"""把特徵設定切換回該 checkpoint 訓練當時的版本。
傳入完整 checkpoint 或 flags dict 皆可。這一步很重要 ——
config 改過之後,舊權重的通道數會對不上。
"""
flags = ckpt_or_flags
if isinstance(ckpt_or_flags, dict) and "state_dict" in ckpt_or_flags:
flags = flags_from_ckpt(ckpt_or_flags)
for k, v in (flags or {}).items():
if hasattr(C, k):
setattr(C, k, v)
# ====================================================================== 資料增強
def limb_scale(coords: np.ndarray, rng: np.random.Generator, amount: float) -> np.ndarray:
"""沿骨架樹隨機縮放每一段肢段的長度,模擬不同體型的人。
做法:由根節點往外走,把每個節點相對父節點的骨向量乘上隨機倍率,
再依新的骨向量重建位置(等同前向運動學)。動作形狀保留,身體比例改變。
"""
parents, order = _parents()
P = coords.shape[1]
factor = 1.0 + rng.uniform(-amount, amount, size=P)
bone = coords - coords[:, np.where(parents >= 0, parents, np.arange(P)), :]
out = coords.copy()
for v in order:
p = parents[v]
if p < 0:
continue
out[:, v] = out[:, p] + bone[:, v] * factor[v]
return out
def _drop_span(coords, mask, rng, node_idx):
"""讓指定節點在一段隨機時間內「偵測失敗」。
模擬真實情形:mask 標 0,座標用該時段前後的值線性內插填補
—— 這正是 normalize_sequence 遇到 NaN 時的處理方式。
"""
T = coords.shape[0]
span = max(2, int(T * rng.uniform(0.1, 0.5)))
start = int(rng.integers(0, max(1, T - span)))
end = min(T, start + span)
if start == 0 and end >= T:
coords[:, node_idx] = 0.0
else:
seg = coords[start:end, node_idx]
a = coords[max(start - 1, 0), node_idx]
b = coords[min(end, T - 1), node_idx]
# ramp 的維度要跟著 node_idx 是單一節點還是一段 slice 調整
ramp = np.linspace(0, 1, end - start).reshape((-1,) + (1,) * (seg.ndim - 1))
coords[start:end, node_idx] = a[None] * (1 - ramp) + b[None] * ramp
mask[start:end, node_idx] = 0.0
return coords, mask
def augment(coords: np.ndarray, mask: np.ndarray, rng: np.random.Generator):
"""在正規化座標空間做幾何 + 時間增強。coords: (T, P, D)"""
T, P, D = coords.shape
# 體型:隨機改變肢段比例(domain randomization 的核心)
if C.AUG_LIMB_SCALE > 0:
coords = limb_scale(coords, rng, C.AUG_LIMB_SCALE)
# 時間:隨機裁掉頭尾,模擬即時視窗沒對齊完整動作
if C.AUG_TEMPORAL_CROP > 0 and T > C.MIN_FRAMES * 2:
keep_ratio = 1.0 - rng.uniform(0, C.AUG_TEMPORAL_CROP)
new_T = max(C.MIN_FRAMES, int(T * keep_ratio))
start = int(rng.integers(0, T - new_T + 1))
coords, mask = coords[start:start + new_T], mask[start:start + new_T]
T = new_T
# 偵測失敗:整隻手短暫消失
if C.AUG_HAND_DROP > 0 and rng.random() < C.AUG_HAND_DROP:
sl = C.LHAND_SLICE if rng.random() < 0.5 else C.RHAND_SLICE
coords, mask = coords.copy(), mask.copy()
coords, mask = _drop_span(coords, mask, rng, slice(sl.start, sl.stop))
# 偵測失敗:零星關節掉點
if C.AUG_JOINT_DROP > 0:
n_drop = int(P * C.AUG_JOINT_DROP * rng.random())
if n_drop > 0:
coords, mask = coords.copy(), mask.copy()
for v in rng.choice(P, size=n_drop, replace=False):
coords, mask = _drop_span(coords, mask, rng, int(v))
# 時間:隨機加減速
if C.AUG_TIME_WARP > 0:
factor = 1.0 + rng.uniform(-C.AUG_TIME_WARP, C.AUG_TIME_WARP)
new_T = max(C.MIN_FRAMES, int(round(T * factor)))
coords = _resample_time(coords, new_T)
mask = _resample_time(mask[:, :, None], new_T)[:, :, 0]
T = new_T
# 時間:隨機丟影格(模擬掉幀)
if C.AUG_FRAME_DROP > 0 and T > C.MIN_FRAMES * 2:
keep = rng.random(T) > C.AUG_FRAME_DROP
if keep.sum() >= C.MIN_FRAMES:
coords, mask = coords[keep], mask[keep]
# 幾何:旋轉(只作用在 x,y 平面)
if C.AUG_ROTATE_DEG > 0 and D >= 2:
th = np.deg2rad(rng.uniform(-C.AUG_ROTATE_DEG, C.AUG_ROTATE_DEG))
c, s = np.cos(th), np.sin(th)
xy = coords[:, :, :2]
coords = coords.copy()
coords[:, :, 0] = xy[:, :, 0] * c - xy[:, :, 1] * s
coords[:, :, 1] = xy[:, :, 0] * s + xy[:, :, 1] * c
# 幾何:縮放 / 平移 / 雜訊
if C.AUG_SCALE > 0:
coords = coords * (1.0 + rng.uniform(-C.AUG_SCALE, C.AUG_SCALE))
if C.AUG_SHIFT > 0:
coords = coords + rng.uniform(-C.AUG_SHIFT, C.AUG_SHIFT, size=(1, 1, D))
if C.AUG_NOISE > 0:
coords = coords + rng.normal(0, C.AUG_NOISE, size=coords.shape)
# 左右鏡像(預設關閉,見 config 說明)
if C.AUG_MIRROR and rng.random() < 0.5 and P == C.NUM_POINTS:
coords = coords.copy()
coords[:, :, 0] *= -1
lh, rh = C.LHAND_SLICE, C.RHAND_SLICE
coords[:, lh], coords[:, rh] = coords[:, rh].copy(), coords[:, lh].copy()
mask = mask.copy()
mask[:, lh], mask[:, rh] = mask[:, rh].copy(), mask[:, lh].copy()
return coords.astype(np.float32), mask.astype(np.float32)
# ====================================================================== Dataset
class SignDataset(Dataset):
"""把 .npy 全部預先載入並前處理好,放在記憶體裡。
前處理(正規化、補值)跟增強無關,所以只做一次;
訓練用的增強在 __getitem__ 才即時套用。
用 subset() 取子集時會共用同一份記憶體,不會重複載入。
"""
def __init__(self, files=None, labels=None, train: bool = False,
seed: int = C.SEED, _items=None, _labels=None, verbose=True):
self.train = train
self.rng = np.random.default_rng(seed)
if _items is not None: # 由 subset() 建立的檢視
self.items, self.labels = _items, _labels
return
self.labels = np.asarray(labels, dtype=np.int64)
self.items = []
n = len(files)
for k, p in enumerate(files):
raw = np.load(p)
if raw.shape[0] < C.MIN_FRAMES:
raw = _resample_time(np.asarray(raw, dtype=np.float32), C.MIN_FRAMES)
self.items.append(normalize_sequence(raw))
if verbose and n > 500 and (k + 1) % 500 == 0:
print(f" 前處理 {k + 1}/{n} …", flush=True)
def subset(self, indices, train: bool, seed: int = C.SEED) -> "SignDataset":
"""共用已前處理好的資料,只換索引與是否增強。"""
idx = list(indices)
return SignDataset(
train=train, seed=seed,
_items=[self.items[i] for i in idx],
_labels=self.labels[idx],
)
def __len__(self):
return len(self.items)
def __getitem__(self, i):
coords, mask = self.items[i]
if self.train:
coords, mask = augment(coords, mask, self.rng)
coords = _resample_time(coords, C.SEQ_LEN)
mask = _resample_time(mask[:, :, None], C.SEQ_LEN)[:, :, 0]
return build_features(coords, mask), self.labels[i]
def stratified_split(labels: np.ndarray, val_ratio: float, seed: int):
"""每個類別各留 val_ratio 當驗證集,至少留 1 筆。"""
rng = np.random.default_rng(seed)
train_idx, val_idx = [], []
for c in np.unique(labels):
idx = np.where(labels == c)[0]
rng.shuffle(idx)
n_val = max(1, int(round(len(idx) * val_ratio))) if len(idx) > 1 else 0
val_idx.extend(idx[:n_val])
train_idx.extend(idx[n_val:])
return np.array(sorted(train_idx)), np.array(sorted(val_idx))
|