refactor: 重构数据管道模块

This commit is contained in:
2026-03-20 21:36:37 +08:00
parent 7045de9c32
commit 3fb2326115
20 changed files with 739 additions and 345 deletions
+41
View File
@@ -0,0 +1,41 @@
from typing import List
import torch
from torch import Tensor
class SequencePacker:
"""序列打包策略"""
def __init__(self, pack_size: int, pad_value: int = 0):
self.pack_size = pack_size
self.pad_value = pad_value
def pack(self, sequences: List[Tensor]) -> List[Tensor]:
"""打包序列到固定大小"""
packages = []
sequences.sort(key=lambda x: x.numel(), reverse=True)
current_pack = torch.full((self.pack_size,), self.pad_value, dtype=torch.int32)
current_pos = 0
for tensor in sequences:
tensor = tensor[:self.pack_size] if tensor.numel() > self.pack_size else tensor
tensor_size = tensor.numel()
if current_pos + tensor_size > self.pack_size:
packages.append(current_pack)
current_pack = torch.full((self.pack_size,), self.pad_value, dtype=torch.int32)
current_pos = 0
current_pack[current_pos:current_pos + tensor_size] = tensor
current_pos += tensor_size
if current_pos > 0:
packages.append(current_pack)
return packages
def pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> List[Tensor]:
"""向后兼容的函数接口"""
return SequencePacker(pack_size, pad_value).pack(sequences)