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

PyTorch实现数据集重叠批次采样及标签偏移的方法

实现带重叠元素的批次与标签偏移

要实现你需求的重叠批次和标签偏移效果,我们可以通过**自定义采样器(Sampler)**控制DataLoader的批次索引逻辑,同时修改Dataset的__getitem__方法支持标签偏移。以下是具体实现:

1. 自定义重叠采样器(OverlapSampler)

这个采样器会生成所有带指定重叠元素的批次索引,核心是用batch_size - overlap_n作为步长,确保下一批复用前一批的最后n个元素:

from torch.utils.data import Sampler
import math

class OverlapSampler(Sampler):
    def __init__(self, data_len, batch_size, overlap_n):
        self.data_len = data_len
        self.batch_size = batch_size
        self.overlap_n = overlap_n
        # 计算总批次数量,最后一批元素不足batch_size时保留
        self.num_batches = math.ceil((data_len - overlap_n) / (batch_size - overlap_n))

    def __iter__(self):
        indices = []
        for i in range(self.num_batches):
            start = i * (self.batch_size - self.overlap_n)
            end = start + self.batch_size
            # 处理最后一批可能超出数据长度的情况,直接截断
            batch_indices = list(range(start, min(end, self.data_len)))
            indices.extend(batch_indices)
        return iter(indices)

    def __len__(self):
        # 返回总采样元素数
        return sum(min(self.batch_size, self.data_len - i*(self.batch_size - self.overlap_n)) for i in range(self.num_batches))

2. 修改Dataset支持标签偏移

给CustomTextDataset添加label_offset参数,让标签取idx + label_offset位置的值:

import torch
from torch.utils.data import Dataset, DataLoader

class CustomTextDataset(Dataset):
    def __init__(self, X, y, label_offset=0):
        self.X = X
        self.y = y
        self.label_offset = label_offset
        # 校验偏移后标签索引是否合法
        assert len(y) - label_offset > 0, "标签偏移超出数据范围"

    def __len__(self):
        # 标签偏移后,有效样本数为原长度减去偏移量
        return len(self.y) - self.label_offset

    def __getitem__(self, idx):
        data = self.X[idx]
        # 取偏移后的标签
        label = self.y[idx + self.label_offset]
        return data, label

3. 组合使用实现需求

以你的示例数据为例,实现重叠1个元素(n=1)、标签偏移m=1的效果:

# 定义数据和标签
X = [1, 2, 3, 4, 5]
y = [0, 0, 1, 0, 1]

# 初始化Dataset,设置标签偏移m=1
td = CustomTextDataset(X, y, label_offset=1)
# 初始化重叠采样器:数据长度为Dataset有效长度,batch_size=2,重叠n=1
sampler = OverlapSampler(data_len=len(td), batch_size=2, overlap_n=1)
# 初始化DataLoader
ddl = DataLoader(td, batch_size=2, sampler=sampler)

# 遍历输出批次
for sample, target in ddl:
    print(f"样本批次: {sample.numpy()}, 标签批次: {target.numpy()}")

输出结果:

样本批次: [1 2], 标签批次: [0 1]
样本批次: [2 3], 标签批次: [1 0]
样本批次: [3 4], 标签批次: [0 1]
样本批次: [4 5], 标签批次: [1]

关键参数说明

  • overlap_n:每批与上一批重叠的元素数量,比如设置为2时,下一批会复用前一批的最后2个元素
  • label_offset:标签相对于数据索引的偏移量m,即取y[idx+m]作为当前数据的标签
  • 如果需要最后一批强制取满batch_size,可以在Sampler中调整逻辑(比如循环补全或丢弃最后一批)

内容的提问来源于stack exchange,提问作者David

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 22:31:16