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

R语言torch框架下无需重定义的数据集合并方法

在R torch中实现自定义数据集的加法合并

实现思路

通过定义通用的ConcatDataset类整合多个数据集,并为自定义feature数据集类重载+运算符,实现dataset_1 + dataset_2的合并效果,无需修改原有数据集的核心逻辑。

完整代码实现

首先添加合并数据集的类定义与运算符重载:

library(torch)

# 定义用于合并数据集的ConcatDataset类
ConcatDataset <- dataset(
  name = "ConcatDataset",
  
  initialize = function(datasets) {
    self$datasets <- datasets
    # 计算各数据集长度及累积长度,快速定位索引所属数据集
    self$lengths <- sapply(datasets, function(ds) ds$.length())
    self$cum_lengths <- cumsum(self$lengths)
  },
  
  .getitem = function(i) {
    # 确定当前索引对应的数据集
    dataset_idx <- findInterval(i - 1, self$cum_lengths) + 1
    # 计算在对应数据集中的局部索引
    local_idx <- if (dataset_idx == 1) {
      i
    } else {
      i - self$cum_lengths[dataset_idx - 1]
    }
    # 从目标数据集获取元素
    self$datasets[[dataset_idx]]$.getitem(local_idx)
  },
  
  .length = function() {
    sum(self$lengths)
  }
)

# 为自定义feature数据集重载+运算符
`+.feature` <- function(a, b) {
  # 检查时间步长一致性,避免维度不匹配
  if (a$n_timesteps != b$n_timesteps) {
    stop("合并的两个数据集必须具有相同的n_timesteps参数")
  }
  ConcatDataset(list(a, b))
}

接着使用原有代码创建数据集并执行合并:

# 原有自定义数据集定义
load_dataset <- dataset(
  name = "feature",
  
  initialize = function(feaute,labels, n_timesteps, sample_frac = 0.5) {
    
    self$n_timesteps <- n_timesteps  
    self$x <- feaute %>% torch_tensor()
    self$y <- labels %>% torch_tensor()
    
    
    n <- nrow(self$x) - self$n_timesteps + 1
    self$starts <- sort(sample.int(
      n = n,
      size = n * sample_frac
    ))
    
  },
  
  .getitem = function(i) {
    
    start <- self$starts[i]
    end <- start + self$n_timesteps - 1
    
    list(
      x = self$x[start:end],
      y = self$y[start]
    )
    
  },
  
  .length = function() {
    length(self$starts)
  }
  
)

# 创建两个数据集实例
x1 = cbind(rnorm(10),rnorm(10))
y1 = rnorm(10)
dataset_1 <- load_dataset(feaute = x1, labels = y1,n_timesteps = 1,sample_frac = 1)

x2 = cbind(rnorm(10),rnorm(10))
y2 = rnorm(10)
dataset_2 <- load_dataset(feaute = x2, labels = y2,n_timesteps = 1,sample_frac = 1)

# 执行合并操作
combine_dataset <- dataset_1 + dataset_2

验证合并效果

通过以下代码确认合并结果:

# 查看合并后数据集总长度(应为10+10=20)
cat("合并后数据集长度:", combine_dataset$.length(), "\n")

# 获取第1个元素(来自dataset_1)
cat("第1个元素:\n")
print(combine_dataset$.getitem(1))

# 获取第11个元素(来自dataset_2)
cat("第11个元素:\n")
print(combine_dataset$.getitem(11))

说明

  • 合并后的数据集保留原数据集的采样逻辑,按顺序返回两个数据集的元素,不会修改原有数据或重新采样。
  • 重载的+运算符会检查n_timesteps参数一致性,避免后续取元素时出现维度错误。
  • 合并后的combine_dataset可直接传入torch::dataloader使用,用法与单个数据集完全一致。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 21:45:01