R中Rescorla-Wagner模型应用问题:无法获取最优学习率α
Rescorla-Wagner模型最优学习率求解:对数似然补全与代码修正
问题背景
需要拟合Rescorla-Wagner模型求解最优学习率α,数据包含ID(被试标识)、Turn_no(试次编号)、p_mean(初始预测值)、t_mean(观测结果)字段,模型更新公式为:p(turn) = p(turn-1) + α × [p(turn-1) - t(turn-1)]
现有函数框架,但对数似然计算部分(ll <- -sum(log(??)))无法完成,同时代码存在变量引用错误、未处理被试分组的问题。
核心问题与修正方案
- 变量引用错误:原代码中
num_reps <- length(df$Turn_no)里的df未定义,应改为输入参数data。 - 未按被试分组:不同被试的学习过程独立,需按
ID分组处理,避免跨被试更新预测值。 - 对数似然计算逻辑:
- 若
t_mean是二元观测值(如0/1,代表强化与否),每个试次的观测概率为:当t_mean=1时取预测值p_pred,t_mean=0时取1-p_pred; - 若
t_mean是连续观测值,需假设残差服从正态分布,用正态密度函数计算概率。
- 若
完整修正代码
二元观测场景(最常见)
RWmodel = function(data, par) { alpha <- par[1] # 按被试ID分组,独立处理每个被试的学习过程 grouped_data <- split(data, data$ID) total_log_likelihood <- 0 for(subj_data in grouped_data) { # 复制初始预测值,避免修改原始数据 p_pred <- subj_data$p_mean t_obs <- subj_data$t_mean num_turns <- nrow(subj_data) # 从第2个试次开始更新预测值 for(i in 2:num_turns) { # 计算预测误差 prediction_error <- p_pred[i-1] - t_obs[i-1] # 更新预测值 p_pred[i] <- p_pred[i-1] + alpha * prediction_error } # 计算单个被试的对数似然 # 给概率加极小值,避免log(0)报错 trial_probs <- ifelse(t_obs == 1, p_pred, 1 - p_pred) trial_probs <- pmax(trial_probs, 1e-10) total_log_likelihood <- total_log_likelihood + sum(log(trial_probs)) } # 返回负对数似然(适配optim等最小化优化函数) return(-total_log_likelihood) }
连续观测场景(需额外估计标准差)
如果t_mean是连续值,需在参数中加入标准差sd,修改后的代码如下:
RWmodel_continuous = function(data, par) { alpha <- par[1] sd_val <- par[2] grouped_data <- split(data, data$ID) total_log_likelihood <- 0 for(subj_data in grouped_data) { p_pred <- subj_data$p_mean t_obs <- subj_data$t_mean num_turns <- nrow(subj_data) for(i in 2:num_turns) { prediction_error <- p_pred[i-1] - t_obs[i-1] p_pred[i] <- p_pred[i-1] + alpha * prediction_error } # 用正态分布密度计算每个试次的概率 trial_probs <- dnorm(t_obs, mean = p_pred, sd = sd_val) trial_probs <- pmax(trial_probs, 1e-10) total_log_likelihood <- total_log_likelihood + sum(log(trial_probs)) } return(-total_log_likelihood) }
调用方式示例
使用optim函数求解最优α:
# 初始参数猜测(二元场景仅需α) init_par <- c(0.5) # 拟合模型 fit <- optim(par = init_par, fn = RWmodel, data = your_data) # 最优学习率 optimal_alpha <- fit$par[1]
内容的提问来源于stack exchange,提问作者LoveDekel
相关产品推荐
相关产品推荐

