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
相关产品推荐
相关产品推荐

