refactor: deduplicate preprocessing kernel and BFD packing

- Extract shared core (mask building, primary-id extraction, tensorisation, position-id generation) to astrai/preprocessing/core.py; Pipeline and TokenizeTransform both consume it, eliminating ~60% duplicated logic
- Promote BFD _plan to module-level plan_bfd(lengths, max_len) returning pure index bins; BFDPacking.apply and evaluate_ifd._pack_bins both call it, removing the second BFD implementation
- Split Pipeline._flush (49 lines) into _inject_doc_reset_position_ids + _inject_continuous_position_ids + _to_tensors; split Pipeline.run by delegating record iteration to core.iter_raw_records
- Remove dead no-op pop/塞回 in Pipeline.run (L110-111)
This commit is contained in:
2026-07-19 12:27:56 +08:00
parent 17127f8b3c
commit 31c22dc043
6 changed files with 297 additions and 146 deletions
+12 -18
View File
@@ -26,28 +26,22 @@ import torch.nn.functional as F
import tqdm
from astrai.model import AutoModel
from astrai.preprocessing.packing import plan_bfd
from astrai.tokenize import AutoTokenizer
def _pack_bins(pairs, max_len):
"""BFD bin packing: pack (c+r) into bins of max total length."""
indexed = sorted(enumerate(pairs), key=lambda x: -(len(x[1][0]) + len(x[1][1])))
bins = []
lengths = []
for orig_idx, (c, r) in indexed:
size = len(c) + len(r)
best_bin = -1
for bi, rem in enumerate(lengths):
if rem >= size:
if best_bin < 0 or rem < lengths[best_bin]:
best_bin = bi
if best_bin >= 0:
bins[best_bin].append((orig_idx, c, r))
lengths[best_bin] -= size
else:
bins.append([(orig_idx, c, r)])
lengths.append(max_len - size)
return bins
"""BFD bin packing: pack (c+r) into bins of max total length.
Reuses :func:`plan_bfd` so the BFD heuristic stays single-sourced.
"""
# Treat each pair as a single sequence of length len(c)+len(r) for
# planning purposes; plan_bfd works on pure lengths.
fake_sequences = [[0] * (len(c) + len(r)) for c, r in pairs]
plan = plan_bfd(fake_sequences, max_len)
return [
[(i, pairs[i][0], pairs[i][1]) for i in bin_indices] for bin_indices in plan
]
def _resolve_sentinel_ids(tokenizer, sentinel_text):