如何基于multi-armed bandit求解未知分布加权骰子的最大平均收益
解决方案:基于汤普森采样的多臂老虎机实现
核心优势
汤普森采样是当前解决这类黑盒最优选择问题效率最高的算法之一,完全适配你的场景需求:
- 自动实现探索与利用的平衡:对表现差的候选骰子只会分配极少的投掷次数,资源会自动向高潜力骰子倾斜,整体效率远高于随机测试、ε-贪心等方案
- 无需提前知道骰子的分布特征,仅依赖实际投掷返回的结果即可运行
- 天然支持混合类型入参:只需要为每组合法入参生成唯一标识即可,不需要对非数值型参数做特殊转换
混合参数处理方案
你可以直接将不同入参按固定顺序拼接为唯一字符串作为骰子的ID,例如类型为lu、效果值为1的骰子ID为lu:1,每个ID对应一个独立的老虎机臂,不需要对非数值入参做额外的编码转换。
JavaScript 完整代码实现
// 汤普森采样多臂老虎机实现 class ThompsonSamplingBandit { constructor() { // 存储每个臂的统计数据:累计收益总和、投掷次数 this.armStats = new Map(); // 先验标准差,可根据你的收益波动范围调整,这里适配1-20的投掷结果 this.priorStd = 5; } // 注册新的骰子入参组合,返回对应的唯一ID registerArm(params) { const armId = params.join(':'); if (!this.armStats.has(armId)) { this.armStats.set(armId, { sum: 0, count: 0 }); } return armId; } // 从正态分布采样随机值(Box-Muller变换) #sampleNormal(mean, std) { const u1 = Math.random(); const u2 = Math.random(); const z = Math.sqrt(-2 * Math.log(u1)) * Math.cos(2 * Math.PI * u2); return mean + z * std; } // 选择当前要投掷的臂ID selectArm() { let maxSample = -Infinity; let selectedArm = null; for (const [armId, stats] of this.armStats.entries()) { // 计算当前臂的后验均值和标准差 const mean = stats.count === 0 ? 10.5 : stats.sum / stats.count; // 1-20的均匀分布均值为10.5作为初始先验 const std = stats.count === 0 ? this.priorStd : this.priorStd / Math.sqrt(stats.count); // 采样一个值 const sample = this.#sampleNormal(mean, std); if (sample > maxSample) { maxSample = sample; selectedArm = armId; } } return selectedArm; } // 更新选中臂的统计数据 updateArm(armId, reward) { const stats = this.armStats.get(armId); stats.sum += reward; stats.count += 1; } // 获取当前最优的臂和对应的平均收益 getBestArm() { let bestMean = -Infinity; let bestArm = null; for (const [armId, stats] of this.armStats.entries()) { const mean = stats.count === 0 ? 0 : stats.sum / stats.count; if (mean > bestMean) { bestMean = mean; bestArm = armId; } } return { armId: bestArm, averageReward: bestMean }; } } // ------------------------------ 示例使用 ------------------------------ // 替换为你实际的骰子投掷逻辑 function rollDie(armId) { const [type, effect] = armId.split(':'); // 示例中的Die-A:lu:1 均值12.4 if (type === 'lu') { const rand = Math.random(); return rand < 0.2 ? 20 : Math.floor(Math.random() * 19) + 1; } // 示例中的Die-B:ul:1 均值8.6 if (type === 'ul') { const rand = Math.random(); return rand < 0.2 ? 1 : Math.floor(Math.random() * 19) + 2; } // 其他骰子逻辑可自行扩展 return Math.floor(Math.random() * 20) + 1; } // 初始化老虎机 const bandit = new ThompsonSamplingBandit(); // 注册所有候选骰子入参组合 bandit.registerArm(['lu', '1']); bandit.registerArm(['ul', '1']); // 可继续注册更多入参组合,例如 bandit.registerArm(['xx', '2']) 等 // 总投掷次数,可根据你的资源情况调整 const totalTrials = 2000; for (let i = 0; i < totalTrials; i++) { const selectedArm = bandit.selectArm(); const reward = rollDie(selectedArm); bandit.updateArm(selectedArm, reward); } // 输出结果 const result = bandit.getBestArm(); console.log('最优骰子ID:', result.armId); console.log('估算平均收益:', result.averageReward.toFixed(2));
使用说明
- 将示例中的
rollDie函数替换为你实际的投掷逻辑即可 - 若候选骰子数量较多,可以适当调大
totalTrials的数值,投掷次数越多,平均收益的估算精度越高 - 如果你的投掷结果的波动范围远大于1-20,可以调整构造函数中的
priorStd参数,设置为和收益波动范围匹配的数值即可提升采样效率
内容的提问来源于stack exchange,提问作者Zlanor
相关产品推荐
相关产品推荐

