如何高效生成n维概率向量网格?求基于numpy的优化实现方案
Hey, I get it—generating that probability simplex grid with brute-force filtering gets really ugly fast when n or v grows. Let's break down a way to do this efficiently with NumPy, skipping all the redundant calculations entirely.
The Core Idea
First, let's reframe the problem: your probability vector elements are evenly spaced from 0 to 1 with v steps, which means each element can be written as k_i/(v-1) where k_i is a non-negative integer. And since the sum of the vector is 1, that translates to k_1 + k_2 + ... + k_n = v-1.
Instead of generating all possible combinations and filtering, we can directly generate all valid integer k tuples that add up to v-1, then normalize them to get our probability vectors. This avoids wasting time on all the invalid points that don't sum to 1.
Efficient NumPy Implementation
Here's a clean, fast implementation using combinatorics (the "stars and bars" method) to generate exactly the points we need:
import numpy as np import itertools from scipy.special import comb def generate_simplex_grid(n, v): m = v - 1 # Integer sum target: 0 to m maps to 0.0 to 1.0 when divided by m if n == 1: return np.array([[1.0]]) # Calculate total number of points (stars and bars formula) total_points = int(comb(m + n - 1, n - 1)) grid = np.zeros((total_points, n), dtype=np.float64) # Generate all ways to split m "stars" into n "bins" using dividers dividers = np.array(list(itertools.combinations(range(m + n - 1), n - 1))) # Compute the count of stars in each bin (k_i values) k_values = np.diff(np.concatenate( (np.zeros((dividers.shape[0], 1), dtype=int), dividers, np.full((dividers.shape[0], 1), m + n - 1)), axis=1)) - 1 # Convert integer counts to probability values grid = k_values / m return grid
How This Works
- Integer Sum Conversion: By working with integers first, we avoid floating-point precision issues (like checking
sum(point) == 1which can fail due to rounding). - Stars and Bars: This combinatorial method generates all valid splits of
munits intonparts directly. For example, with n=3, m=2 (v=3), we're splitting 2 units into 3 bins—exactly the 6 valid combinations we need. - No Redundancy: We only generate the points that actually lie on the simplex, so there's no filtering step wasting CPU cycles.
Test It Out
For your example (n=3, v=3):
grid = generate_simplex_grid(3, 3) print(grid)
You'll get exactly the same result as your original code, but without generating all 27 redundant points first:
[[0. 0. 1. ] [0. 0.5 0.5] [0. 1. 0. ] [0.5 0. 0.5] [0.5 0.5 0. ] [1. 0. 0. ]]
Notes
- Combination Explosion: Keep in mind that the number of points grows combinatorially—this is inherent to the simplex grid, not the method. But this implementation is as efficient as possible because it generates only the necessary points.
- Scipy Dependency: If you don't want to use
scipy.special.comb, you can calculate the combination manually or just get the length of thedividerslist to initialize the grid (though calculating upfront is more memory-efficient for large cases).
内容的提问来源于stack exchange,提问作者Schiphol

