R中torch与targets包兼容问题:无法读取dataset类目标
解决R中torch自定义Dataset与targets包的互操作性问题
问题根源
torch的torch_tensor对象依赖底层C++外部指针,直接用targets默认的序列化方式(如rds格式)保存自定义Dataset实例时,只会记录指针地址而非实际数据。当你用tar_read()读取时,原指针指向的内存已释放,就会触发external pointer is not valid错误。而format="torch"或tar_torch仅对torch原生对象(如模型、单个张量)提供支持,无法自动处理包含多张量的自定义Dataset类。
解决方案
方案1:分离预处理与Dataset初始化(推荐)
把数据预处理步骤拆为独立target,保存预处理后的张量数据;再定义动态生成Dataset的target,依赖预处理好的数据。这种方式更贴合targets的工作流,避免序列化外部指针的问题。
修改后的代码
library(torch) library(targets) library(dplyr) library(tidymodels) # 独立的预处理函数,返回预处理后的张量集合 prepare_nn_data <- function(df) { target_col <- df$claim_ind_cov_1_2_3_4_5_6 %>% as.integer() %>% `-`(1) %>% as.matrix() tele_cols <- df %>% select(starts_with(c("h_", "p_", "vmo", "vma"))) %>% as.matrix() class_df <- select(df, expo:years_licensed, distance) rec_class <- recipe(~ ., data = class_df) %>% step_impute_median(commute_distance, years_claim_free) %>% step_other(all_nominal(), threshold = 0.05) %>% step_dummy(all_nominal()) %>% prep() class_cols <- juice(rec_class) %>% as.matrix() list( x = list( tele = torch_tensor(tele_cols), class = torch_tensor(class_cols) ), y = torch_tensor(target_col) ) } # 调整Dataset定义,改为接收预处理好的数据 nn_dataset <- dataset( name = "nn_dataset", initialize = function(data) { self$tele <- data$x$tele self$class <- data$x$class self$y <- data$y }, .getitem = function(i) { list( x = list( tele = self$tele[i, ], class = self$class[i, ] ), y = self$y[i, ] ) }, .length = function() { self$y$size()[[1]] } )
定义targets
# 保存预处理后的张量数据,用torch格式确保序列化有效 tar_target( name = valid_nn_data, command = prepare_nn_data(valid_df), format = "torch" ) # 动态生成Dataset,依赖已加载的有效张量数据 tar_target( name = valid_nn_dataset, command = nn_dataset(valid_nn_data), format = "rds" )
此时执行tar_read(valid_nn_dataset),会先加载预处理好的有效张量,再初始化Dataset,不会出现指针失效问题。
方案2:自定义Dataset的序列化方法
如果需要直接保存Dataset实例,可以给自定义Dataset类添加serialize和deserialize方法,让targets能正确处理张量数据的序列化与重建。
修改后的Dataset定义
nn_dataset <- dataset( name = "nn_dataset", initialize = function(df) { # 支持从原始数据或序列化后的张量集合初始化 if (inherits(df, "list") && all(c("tele", "class", "y") %in% names(df))) { self$tele <- df$tele self$class <- df$class self$y <- df$y } else { data <- self$prepare_data(df) self$tele <- data$x$tele self$class <- data$x$class self$y <- data$y } }, .getitem = function(i) { list( x = list( tele = self$tele[i, ], class = self$class[i, ] ), y = self$y[i, ] ) }, .length = function() { self$y$size()[[1]] }, prepare_data = function(df) { # 原预处理逻辑保持不变 target_col <- df$claim_ind_cov_1_2_3_4_5_6 %>% as.integer() %>% `-`(1) %>% as.matrix() tele_cols <- df %>% select(starts_with(c("h_", "p_", "vmo", "vma"))) %>% as.matrix() class_df <- select(df, expo:years_licensed, distance) rec_class <- recipe(~ ., data = class_df) %>% step_impute_median(commute_distance, years_claim_free) %>% step_other(all_nominal(), threshold = 0.05) %>% step_dummy(all_nominal()) %>% prep() class_cols <- juice(rec_class) %>% as.matrix() list( x = list( tele = torch_tensor(tele_cols), class = torch_tensor(class_cols) ), y = torch_tensor(target_col) ) }, # 自定义序列化:提取需要保存的张量数据 serialize = function() { list( tele = self$tele, class = self$class, y = self$y ) }, # 自定义反序列化:从张量数据重建Dataset deserialize = function(data) { self$initialize(data) } )
定义target
tar_target( name = target_name, command = nn_dataset(valid_df), format = "torch" )
targets会调用Dataset的serialize方法保存张量数据,读取时通过deserialize重建实例,确保张量指针有效。
内容的提问来源于stack exchange,提问作者Francis Duval
相关产品推荐
相关产品推荐

