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

Keyatm协变量模型:四类国家文档-主题分布预测均值一致问题

问题:KeyATM模型预测不同国家组别的文档-主题分布均值完全相同

我将国家划分为LI、LMI、UMI、HI四个组别,原本预期不同组别的文档-主题分布预测均值会存在差异,但运行KeyATM模型后得到的结果完全一致。我参照KeyATM用户指南编写并运行了如下代码:

strata_LI <- by_strata_DocTopic(out, by_var = "CountriesLI",labels = 
  c("Non-LI", "LI")
  )
strata_LMI <- by_strata_DocTopic(out, by_var = "CountriesLMI",labels = 
  c("Non-LMI", "LMI")
  )
strata_UMI <- by_strata_DocTopic(out, by_var = "CountriesUMI",labels = 
  c("Non-UMI", "UMI")
  )
strata_HI <- by_strata_DocTopic(out, by_var = "CountriesHI",labels = 
  c("Non-HI", "HI")
   )

est_LI <- summary(strata_LI)
est_LMI <- summary(strata_LMI)
est_UMI <- summary(strata_UMI)
est_HI <- summary(strata_HI)

new_data <- covariates_get(out)new_data
new_data[, "CountriesLI"] <- 0
new_data[, "CountriesLMI"] <- 0
new_data[, "CountriesUMI"]<- 0
new_data[, "CountriesHI"]<- 0
pred <- predict(out, new_data, label = "Group")

res <- bind_rows(est_LI, est_LMI, est_UMI,est_HI, pred)

pred_LI <- predict(out, new_data, label = "CountriesLI")
pred_LMI <- predict(out, new_data, label = "CountriesLMI")
pred_UMI <- predict(out, new_data, label = "CountriesUMI")
pred_HI <- predict(out, new_data, label = "CountriesHI")

res <- bind_rows(est_LI, est_LMI, est_UMI, est_HI, pred_LI, pred_LMI, 
pred_UMI, pred_HI, pred) 
%>%filter(TopicID %in% c(1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13))

labels <- unique(res$label)
library(ggplot2)
ggplot(res, aes(x = label, ymin = Lower, ymax = Upper, group = Topic)) 
+geom_errorbar(width = 0.1) 
+coord_flip() 
+facet_wrap(~Topic) + 
geom_point(aes(x = label, y = Point)) 
+scale_x_discrete(limits = rev(labels))
+xlab("Countries") 
+ylab(expression(paste("Mean of ", theta))) + 
theme_bw()

核心问题与修正方案

1. 预测数据集构造错误

你在new_data中将所有分组变量设为0,导致后续predict调用时,所有组别预测的都是"非对应组别"的结果,自然输出完全相同。正确做法是为每个组别单独构造数据集,将对应分组变量设为1,其余变量保留原始合理值。

2. 修正后的代码

# 按组别分层计算文档-主题分布
strata_LI <- by_strata_DocTopic(out, by_var = "CountriesLI", labels = c("Non-LI", "LI"))
strata_LMI <- by_strata_DocTopic(out, by_var = "CountriesLMI", labels = c("Non-LMI", "LMI"))
strata_UMI <- by_strata_DocTopic(out, by_var = "CountriesUMI", labels = c("Non-UMI", "UMI"))
strata_HI <- by_strata_DocTopic(out, by_var = "CountriesHI", labels = c("Non-HI", "HI"))

# 获取分层统计结果
est_LI <- summary(strata_LI)
est_LMI <- summary(strata_LMI)
est_UMI <- summary(strata_UMI)
est_HI <- summary(strata_HI)

# 基于原始协变量数据,为每个组别单独构造预测集
new_data <- covariates_get(out)

# 构造LI组预测集:仅将CountriesLI设为1
new_data_LI <- new_data
new_data_LI$CountriesLI <- 1
pred_LI <- predict(out, new_data_LI, label = "LI")

# 构造LMI组预测集
new_data_LMI <- new_data
new_data_LMI$CountriesLMI <- 1
pred_LMI <- predict(out, new_data_LMI, label = "LMI")

# 构造UMI组预测集
new_data_UMI <- new_data
new_data_UMI$CountriesUMI <- 1
pred_UMI <- predict(out, new_data_UMI, label = "UMI")

# 构造HI组预测集
new_data_HI <- new_data
new_data_HI$CountriesHI <- 1
pred_HI <- predict(out, new_data_HI, label = "HI")

# 合并结果并筛选目标主题
res <- bind_rows(est_LI, est_LMI, est_UMI, est_HI, pred_LI, pred_LMI, pred_UMI, pred_HI) %>%
  filter(TopicID %in% 1:13)

# 可视化
library(ggplot2)
labels <- unique(res$label)
ggplot(res, aes(x = label, ymin = Lower, ymax = Upper, group = Topic)) +
  geom_errorbar(width = 0.1) +
  coord_flip() +
  facet_wrap(~Topic) +
  geom_point(aes(x = label, y = Point)) +
  scale_x_discrete(limits = rev(labels)) +
  xlab("国家组别") +
  ylab(expression(paste("Mean of ", theta))) +
  theme_bw()

3. 额外检查项

  • 确认原始数据中四个分组变量的编码正确,每个国家仅属于一个组别,无重叠或缺失值。
  • 检查KeyATM模型训练时是否将这些分组变量作为协变量纳入模型,若模型未引入这些变量,预测结果不会体现组别差异。

内容的提问来源于stack exchange,提问作者topicmodler

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 01:52:38