From 6d52bae1cac7addb5c65b43f3500a9de49b950f3 Mon Sep 17 00:00:00 2001 From: ViperEkura <3081035982@qq.com> Date: Mon, 21 Jul 2025 19:07:16 +0800 Subject: [PATCH] =?UTF-8?q?feat(utils):=20=E6=B7=BB=E5=8A=A0=E5=BA=8F?= =?UTF-8?q?=E5=88=97=E6=89=93=E5=8C=85=E5=8A=9F=E8=83=BD?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- utils.py | 39 +++++++++++++++++++++++++++++++++++++++ 1 file changed, 39 insertions(+) diff --git a/utils.py b/utils.py index b429916..ba69b03 100644 --- a/utils.py +++ b/utils.py @@ -27,6 +27,45 @@ def comprehensive_normalization(text): pattern = re.compile('|'.join(re.escape(k) for k in replacements)) return pattern.sub(lambda m: replacements[m.group()], text) + +def pack_sequences(sequences: List[Tensor], pack_size: int, pad_value: int) -> List[Tensor]: + packages = [] + sequences.sort(key=lambda x: x.numel(), reverse=True) + current_pack = torch.tensor([], dtype=torch.int32) + + for tensor in sequences: + if tensor.numel() > pack_size: + packages.append(tensor[:pack_size]) + continue + + remaining = pack_size - current_pack.numel() + + if remaining == 0: + packages.append(current_pack) + current_pack = tensor + elif tensor.numel() <= remaining: + current_pack = torch.cat([current_pack, tensor]) + else: + padding = torch.full((remaining,), pad_value, dtype=torch.int32) + current_pack = torch.cat([current_pack, padding]) + packages.append(current_pack) + current_pack = tensor + + if current_pack.numel() > 0: + if current_pack.numel() < pack_size: + padding = torch.full( + (pack_size - current_pack.numel(),), + pad_value, + dtype=torch.int32 + ) + current_pack = torch.cat([current_pack, padding]) + else: + current_pack = current_pack[:pack_size] + + packages.append(current_pack) + + return packages + def dump_pkl_files( files: List[str], base_out_dir: str,