You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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 > 0 to parallelize data loading (but don't overdo it—start with 4-8, depending on your CPU cores).
  • Add pin_memory=True if you're using a GPU; this speeds up data transfer from CPU to GPU.
  • Use a custom collate_fn to 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.open with Python 3's built-in open (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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.20 10:38:13