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

Golang中是否存在等价于numpy.random.choice的概率洗牌函数?

Golang中实现类似numpy.random.choice的带权重洗牌方案

Golang标准库的rand.Shuffle确实不支持基于元素权重的洗牌逻辑,gonum包中也没有直接等价于numpy.random.choice(无放回带权重采样生成排列)的函数。不过你可以通过以下方式实现需求:

一、自行实现轮盘赌采样(适合中小规模数据)

轮盘赌算法逻辑直观,适合元素数量不多的场景,核心是每次根据剩余元素的权重总和随机选择下一个元素,直到所有元素都被选入结果。

package main

import (
    "math/rand"
    "time"
)

// WeightedShuffle 对任意类型切片执行带权重洗牌,weights需与elements长度一致
func WeightedShuffle[T any](elements []T, weights []float64) []T {
    // 实际生产环境建议使用固定种子或更安全的随机源,如crypto/rand
    rng := rand.New(rand.NewSource(time.Now().UnixNano()))

    // 初始化元素-权重映射
    type weightedItem struct {
        elem   T
        weight float64
    }
    items := make([]weightedItem, len(elements))
    for i := range elements {
        items[i] = weightedItem{elem: elements[i], weight: weights[i]}
    }

    shuffled := make([]T, 0, len(elements))
    totalWeight := 0.0
    for _, w := range weights {
        totalWeight += w
    }

    for len(items) > 0 {
        // 生成0到当前总权重区间的随机数
        randVal := rng.Float64() * totalWeight
        currentSum := 0.0
        selectedIdx := -1

        // 定位选中的元素
        for i, item := range items {
            currentSum += item.weight
            if currentSum >= randVal {
                selectedIdx = i
                break
            }
        }

        // 将选中元素加入结果,更新剩余元素和总权重
        shuffled = append(shuffled, items[selectedIdx].elem)
        totalWeight -= items[selectedIdx].weight
        items = append(items[:selectedIdx], items[selectedIdx+1:]...)
    }

    return shuffled
}

// 示例调用
func main() {
    elems := []string{"apple", "banana", "cherry", "date"}
    weights := []float64{0.1, 0.2, 0.3, 0.4} // date被选中的概率最高

    result := WeightedShuffle(elems, weights)
    println("带权重洗牌结果:", result)
}

二、利用gonum的分类分布实现(适合规范化概率场景)

如果你的权重已经归一化为总和为1的概率,可以使用gonum的distuv.Categorical类型来实现无放回采样,本质是重复生成符合概率分布的索引,直到所有元素都被选中:

package main

import (
    "math/rand"
    "time"

    "gonum.org/v1/gonum/stat/distuv"
)

func WeightedShuffleWithGonum[T any](elements []T, probs []float64) []T {
    rng := rand.New(rand.NewSource(time.Now().UnixNano()))
    // 初始化分类分布,probs需为归一化后的概率(总和为1)
    catDist := distuv.Categorical{
        Prob: probs,
        Src:  rng,
    }

    used := make([]bool, len(elements))
    shuffled := make([]T, 0, len(elements))

    for len(shuffled) < len(elements) {
        idx := int(catDist.Rand())
        if !used[idx] {
            used[idx] = true
            shuffled = append(shuffled, elements[idx])
        }
    }

    return shuffled
}

三、大规模数据优化:别名采样算法

如果需要处理大规模切片,轮盘赌的O(n)单次采样效率会不足,推荐实现别名采样算法(Alias Method),将单次采样的时间复杂度降至O(1),适合高频调用的场景。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 16:07:04