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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 21:34:59