stzhao's picture
download
raw
6.74 kB
import os
import glob
import numpy as np
import torch
from torch.utils.data import Dataset, DataLoader
from typing import Optional, Tuple, List, Union
class VideoTrajectoryDataset(Dataset):
"""
从视频特征文件中加载轨迹数据,支持两种模式:
- mode='sequence': 每个样本为一个视频的完整特征序列 (features, timesteps)
- mode='derivative': 每个样本为一个时间点 (x_t, t, u),其中 u 由相邻帧差分得到
"""
def __init__(
self,
feature_dir: str,
feature_type: str = 'patch_flatten',
pattern: str = '*_{}.npz',
mode: str = 'derivative', # 'sequence' 或 'derivative'
cache: bool = True,
min_frames: int = 2,
derivative_method: str = 'finite_difference', # 仅 derivative 模式有效
eps: float = 1e-8,
):
"""
Args:
feature_dir: 特征文件目录
feature_type: 特征类型 ('cls' 或 'patch_flatten')
pattern: 文件名匹配模板
mode: 数据返回模式
cache: 是否缓存所有数据到内存
min_frames: 视频最少帧数
derivative_method: 导数计算方法,目前仅支持 'finite_difference'
eps: 防止除以零的小常数
"""
self.feature_dir = feature_dir
self.feature_type = feature_type
self.mode = mode
self.cache = cache
self.min_frames = min_frames
self.derivative_method = derivative_method
self.eps = eps
# 查找特征文件
file_pattern = pattern.format(feature_type)
self.file_list = glob.glob(os.path.join(feature_dir, file_pattern))
if not self.file_list:
raise RuntimeError(f"在 {feature_dir} 中没有找到匹配 {file_pattern} 的文件")
# 加载数据(根据 mode 组织样本列表)
self.samples = [] # 每个元素为 (features, timesteps) 或 (x_t, t, u)
self._load_data()
def _load_data(self):
for fpath in self.file_list:
data = np.load(fpath)
features = data['features'] # (num_frames, feat_dim)
timesteps = data['timesteps'] # (num_frames,)
num_frames = len(features)
if num_frames < self.min_frames:
continue
if self.mode == 'sequence':
# 直接保存整个序列
self.samples.append((features, timesteps))
elif self.mode == 'derivative':
# 为每一对相邻帧构造一个导数样本
for i in range(num_frames - 1):
x_i = features[i]
x_j = features[i+1]
t_i = timesteps[i]
t_j = timesteps[i+1]
dt = t_j - t_i
if dt > self.eps:
u = (x_j - x_i) / dt
else:
# 时间间隔太小,用零导数(或跳过)
u = np.zeros_like(x_i)
# 这里使用前一个时刻作为 x_t,t = t_i
self.samples.append((x_i, t_i, u))
else:
raise ValueError("mode 必须是 'sequence' 或 'derivative'")
if not self.samples:
raise RuntimeError("没有满足条件的样本")
def __len__(self):
return len(self.samples)
def __getitem__(self, idx):
sample = self.samples[idx]
if self.mode == 'sequence':
features, timesteps = sample
return (
torch.from_numpy(features).float(),
torch.from_numpy(timesteps).float()
)
else: # derivative
x_t, t, u = sample
return (
torch.from_numpy(x_t).float(),
torch.tensor(t, dtype=torch.float32),
torch.from_numpy(u).float()
)
def collate_fn_sequence(batch):
"""
用于 'sequence' 模式的 collate 函数。
由于序列长度可能不同,需要填充(padding)并返回 mask。
"""
features_list, timesteps_list = zip(*batch)
lengths = [f.shape[0] for f in features_list]
max_len = max(lengths)
feat_dim = features_list[0].shape[1]
padded_features = torch.zeros(len(batch), max_len, feat_dim)
padded_timesteps = torch.zeros(len(batch), max_len)
mask = torch.zeros(len(batch), max_len, dtype=torch.bool)
for i, (f, t) in enumerate(zip(features_list, timesteps_list)):
seq_len = f.shape[0]
padded_features[i, :seq_len] = f
padded_timesteps[i, :seq_len] = t
mask[i, :seq_len] = True
return padded_features, padded_timesteps, mask
def collate_fn_derivative(batch):
"""
用于 'derivative' 模式的 collate 函数:直接堆叠样本。
"""
x_t_list, t_list, u_list = zip(*batch)
x_t = torch.stack(x_t_list, dim=0)
t = torch.stack(t_list, dim=0)
u = torch.stack(u_list, dim=0)
return x_t, t, u
def create_dataloader(
feature_dir: str,
feature_type: str = 'patch_flatten',
mode: str = 'derivative',
batch_size: int = 1,
shuffle: bool = True,
num_workers: int = 0,
cache: bool = True,
min_frames: int = 2,
**kwargs
) -> DataLoader:
"""
创建 DataLoader 的便捷函数。
"""
dataset = VideoTrajectoryDataset(
feature_dir=feature_dir,
feature_type=feature_type,
mode=mode,
cache=cache,
min_frames=min_frames
)
print("dataset created.")
if mode == 'sequence':
collate_fn = collate_fn_sequence
else:
collate_fn = collate_fn_derivative
dataloader = DataLoader(
dataset,
batch_size=batch_size,
shuffle=shuffle,
num_workers=num_workers,
collate_fn=collate_fn,
**kwargs
)
return dataloader
if __name__ == "__main__":
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
dataloader = create_dataloader(
feature_dir='./video_features/dinov2-small',
feature_type='patch_flatten',
mode='derivative',
batch_size=1,
shuffle=True,
num_workers=0,
cache=True
)
for x_t, t, u in dataloader:
x_t = x_t.to(device)
t = t.to(device)
u = u.to(device)
print(f"shape of x_t: {x_t.shape}")
print(f"t: {t}")
print(f"shape of u: {u.shape}")
# v_pred = model(x_t, t) # 模型预测的向量场
# loss = F.mse_loss(v_pred, u)
# optimizer.zero_grad()
# loss.backward()
# optimizer.step()

Xet Storage Details

Size:
6.74 kB
·
Xet hash:
e6dcac2b53f22a848af58f6e3265783d7c9a3d89b4ed741b4691f1f034c506fe

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.