基于PyTorch加载超大文本文件用于神经机器翻译的技术问询
Hey there! Let's dive into optimizing your LoadUniModal class for those massive 93M-line text files you're working with for NMT. Loading the entire dataset into memory upfront is going to hit you hard with that 8GB file—here are the most practical fixes and optimizations to handle this smoothly:
1. Avoid Loading All Data into Memory Upfront
Your current approach stores all source and target lines in class-level lists, which will eat up gigabytes of RAM (and likely crash your process). Instead, use lazy loading: only read the specific line you need when __getitem__ is called. To make this fast, precompute line offsets so you can jump directly to any line without scanning the whole file.
Here's how to refactor your class:
from torch.utils.data import Dataset import codecs class LoadUniModal(Dataset): def __init__(self, src_path, trg_path, src_vocab, trg_vocab): self.src_path = src_path self.trg_path = trg_path self.src_vocab = src_vocab self.trg_vocab = trg_vocab # Precompute byte offsets for every line (fast random access) self.src_offsets = self._calculate_line_offsets(src_path) self.trg_offsets = self._calculate_line_offsets(trg_path) # Ensure source/target files have matching line counts assert len(self.src_offsets) == len(self.trg_offsets), "Source and target files must have the same number of lines" def _calculate_line_offsets(self, file_path): offsets = [0] with codecs.open(file_path, encoding="utf-8") as f: # Read each line to track the end position of the previous line while f.readline(): offsets.append(f.tell()) return offsets def __len__(self): # Number of lines is total offsets minus the initial 0 return len(self.src_offsets) - 1 def __getitem__(self, idx): # Read source line directly using the precomputed offset with codecs.open(self.src_path, encoding="utf-8") as f: f.seek(self.src_offsets[idx]) src_line = f.readline().strip() # Read corresponding target line with codecs.open(self.trg_path, encoding="utf-8") as f: f.seek(self.trg_offsets[idx]) trg_line = f.readline().strip() # Tokenize and convert to indices src_indices = self.src_vocab.tokenize(src_line) trg_indices = self.trg_vocab.tokenize(trg_line) return src_indices, trg_indices
This way, you only store a small list of integers (the offsets) instead of the entire 8GB text. The tradeoff is a tiny overhead per __getitem__ call, which is negligible compared to the memory savings.
2. Use IterableDataset for Streaming (Even Larger Datasets)
If your dataset is so big that even storing line offsets feels unwieldy, switch to IterableDataset. This lets you stream data line-by-line without loading any metadata upfront, which is perfect for truly massive files. You can even add buffered shuffling to maintain randomness during training.
Example implementation:
from torch.utils.data import IterableDataset import random class LoadUniModalIterable(IterableDataset): def __init__(self, src_path, trg_path, src_vocab, trg_vocab, shuffle=False, buffer_size=10000): self.src_path = src_path self.trg_path = trg_path self.src_vocab = src_vocab self.trg_vocab = trg_vocab self.shuffle = shuffle self.buffer_size = buffer_size # For shuffling small chunks in memory def __iter__(self): worker_info = torch.utils.data.get_worker_info() src_file = codecs.open(self.src_path, encoding="utf-8") trg_file = codecs.open(self.trg_path, encoding="utf-8") # Split workload across DataLoader workers (if using num_workers > 0) if worker_info is not None: num_workers = worker_info.num_workers worker_id = worker_info.id # Skip lines meant for other workers for _ in range(worker_id): src_file.readline() trg_file.readline() # Iterate in steps of num_workers while True: src_line = src_file.readline() trg_line = trg_file.readline() if not src_line or not trg_line: break yield self._process_line(src_line.strip(), trg_line.strip()) # Skip lines for other workers for _ in range(num_workers - 1): src_file.readline() trg_file.readline() else: # Single worker: iterate all lines for src_line, trg_line in zip(src_file, trg_file): yield self._process_line(src_line.strip(), trg_line.strip()) src_file.close() trg_file.close() def _process_line(self, src_line, trg_line): if self.shuffle: # Use a buffer to shuffle small batches of data (avoids loading everything) buffer = [(src_line, trg_line)] while len(buffer) < self.buffer_size: next_src = src_file.readline() next_trg = trg_file.readline() if not next_src or not next_trg: break buffer.append((next_src.strip(), next_trg.strip())) random.shuffle(buffer) for s, t in buffer: yield self.src_vocab.tokenize(s), self.trg_vocab.tokenize(t) else: return self.src_vocab.tokenize(src_line), self.trg_vocab.tokenize(trg_line)
3. Preprocess Offline to Speed Up Training
Tokenizing 93M lines on-the-fly every training run is slow. Instead, preprocess your data once and save it in a binary format (like PyTorch .pt files) that loads quickly. This saves you from repeating the same tokenization work over and over.
Here's a quick preprocessing script:
import os import torch def preprocess_dataset(src_path, trg_path, src_vocab, trg_vocab, save_dir, chunk_size=1_000_000): os.makedirs(save_dir, exist_ok=True) src_chunk = [] trg_chunk = [] chunk_idx = 0 with codecs.open(src_path, encoding="utf-8") as src_f, codecs.open(trg_path, encoding="utf-8") as trg_f: for src_line, trg_line in zip(src_f, trg_f): src_idx = torch.tensor(src_vocab.tokenize(src_line.strip())) trg_idx = torch.tensor(trg_vocab.tokenize(trg_line.strip())) src_chunk.append(src_idx) trg_chunk.append(trg_idx) # Save chunk when we hit the desired size if len(src_chunk) >= chunk_size: torch.save({ "src": src_chunk, "trg": trg_chunk }, os.path.join(save_dir, f"chunk_{chunk_idx}.pt")) src_chunk = [] trg_chunk = [] chunk_idx += 1 # Save any remaining data if src_chunk: torch.save({ "src": src_chunk, "trg": trg_chunk }, os.path.join(save_dir, f"chunk_{chunk_idx}.pt")) # Usage: # preprocess_dataset("src.txt", "trg.txt", src_vocab, trg_vocab, "./preprocessed_data")
Then create a Dataset to load these preprocessed chunks:
class PreprocessedUniModal(Dataset): def __init__(self, preprocessed_dir): self.data = [] # Load all chunks into memory (or load lazily if still too big) for chunk_file in os.listdir(preprocessed_dir): if chunk_file.endswith(".pt"): chunk = torch.load(os.path.join(preprocessed_dir, chunk_file)) self.data.extend(zip(chunk["src"], chunk["trg"])) def __len__(self): return len(self.data) def __getitem__(self, idx): return self.data[idx]
4. Optimize DataLoader for Large Datasets
Tweak your DataLoader settings to avoid bottlenecks:
- Use
num_workers > 0to parallelize data loading (but don't overdo it—start with 4-8, depending on your CPU cores). - Add
pin_memory=Trueif you're using a GPU; this speeds up data transfer from CPU to GPU. - Use a custom
collate_fnto handle dynamic padding (don't pad all sequences upfront—only pad per batch):
def collate_batch(batch): src_tensors, trg_tensors = zip(*batch) # Pad source sequences to the longest in the batch src_padded = torch.nn.utils.rnn.pad_sequence(src_tensors, batch_first=True, padding_value=src_vocab.pad_idx) # Pad target sequences similarly trg_padded = torch.nn.utils.rnn.pad_sequence(trg_tensors, batch_first=True, padding_value=trg_vocab.pad_idx) return src_padded, trg_padded # Initialize DataLoader dataloader = torch.utils.data.DataLoader( dataset, batch_size=32, collate_fn=collate_batch, num_workers=4, pin_memory=True, shuffle=True # Only use this if you're using the offset-based Dataset; for IterableDataset, use buffered shuffling )
5. Quick Wins for Faster File Reading
- Replace
codecs.openwith Python 3's built-inopen(it supports UTF-8 by default and is faster). - For even faster reads, use memory mapping with
mmap—this lets the OS handle file caching efficiently:
import mmap def _calculate_line_offsets_mmap(file_path): offsets = [0] with open(file_path, "rb") as f: with mmap.mmap(f.fileno(), length=0, access=mmap.ACCESS_READ) as mm: while mm.readline(): offsets.append(mm.tell()) return offsets
内容的提问来源于stack exchange,提问作者yusuf

