如何在R包JuliaConnectoR中传递Julia类型适配EvoTrees?
问题说明
使用R的JuliaConnectoR包调用Julia的EvoTrees包时,创建EvoTreeRegressor默认采用Float32类型,但训练数据是R的double(对应Julia的Float64),调用fit_evotree时触发MethodError,核心原因是模型类型与数据类型不匹配。
示例代码
n_obs = 100 n_features = 100 nrounds = 10L set.seed(20221224) x_train = matrix(ncol= n_features, rnorm(n_obs*n_features)) y_train = rnorm(n_obs) library(JuliaConnectoR) evoTrees = juliaImport("EvoTrees") params_evo = evoTrees$EvoTreeRegressor(nrounds=nrounds , loss = as.symbol("linear"), alpha=0.5,lambda=0.0,gamma=0.0) evoTrees_model = evoTrees$fit_evotree(params_evo, x_train, y_train, print_every_n = 50L)
错误信息
Error: Evaluation in Julia failed. Original Julia error message:
MethodError: no method matching fit_evotree(::EvoTrees.EvoTreeRegressor{EvoTrees.Linear, Float32}, ::Matrix{Float64}, ::Vector{Float64}; print_every_n=50) Closest candidates are: fit_evotree(::Union{EvoTrees.EvoTreeClassifier{L, T}, EvoTrees.EvoTreeCount{L, T}, EvoTrees.EvoTreeGaussian{L, T}, EvoTrees.EvoTreeMLE{L, T}, EvoTrees.EvoTreeRegressor{L, T}}; x_train, y_train, w_train, offset_train, x_eval, y_eval, w_eval, offset_eval, metric, early_stopping_rounds, print_every_n, verbosity, fnames, return_logger) where {L, T} at C:\Users\rwarn.julia\packages\EvoTrees\ayRL8\src\fit.jl:309
解决方案
修改要点
- 为
EvoTreeRegressor指定Float64类型:通过关键字参数T传递Julia的Float64类型对象,使用juliaEval("Float64")获取该类型。 - 修正loss参数传递:原代码中
as.symbol("linear")传递的是R符号,EvoTrees需要的是Linear()实例,直接调用evoTrees$Linear()即可。
修改后代码
n_obs = 100 n_features = 100 nrounds = 10L set.seed(20221224) x_train = matrix(ncol= n_features, rnorm(n_obs*n_features)) y_train = rnorm(n_obs) library(JuliaConnectoR) evoTrees = juliaImport("EvoTrees") params_evo = evoTrees$EvoTreeRegressor( nrounds=nrounds, loss=evoTrees$Linear(), alpha=0.5, lambda=0.0, gamma=0.0, T=juliaEval("Float64") ) evoTrees_model = evoTrees$fit_evotree(params_evo, x_train, y_train, print_every_n = 50L)
说明
juliaEval("Float64")用于获取Julia原生的Float64类型,确保模型内部使用的数值类型与训练数据的Float64一致,解决类型不匹配问题。- 调用
evoTrees$Linear()直接获取EvoTrees定义的Linear损失函数实例,符合Julia端的参数要求。
内容的提问来源于stack exchange,提问作者Richi W

