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

如何高效实现12000个客户的状态向量与对应时序多维矩阵的批量乘积运算

Hey there! Let's figure out how to batch-process those 12k customers efficiently instead of looping through each one individually—NumPy's vectorization and smart grouping will be our best friends here.

Efficient Batch Processing for Customer State Vector Matrix Multiplication

First, let's recap the problem: we have 12k customers, each starting with a state vector tied to their current state and amount, and we need to sequentially multiply their vector by a series of year-specific matrices starting from their start year's next matrix up to the 14th index (14th year) matrix.

Step 1: Convert Customer Data to Vectorized Format

First, we'll turn our customer data into a structured NumPy array and map each state (A-J) to an index (0-9) to build our initial state vectors in bulk.

import numpy as np

# Example customer data (replace with your actual dataset)
customer_data = np.array([
    ('ax111', 'A', 3, 300),
    ('ax112', 'D', 4, 4890),
    ('ax113', 'G', 9, 624)
], dtype=[('customer', 'U10'), ('current_state', 'U1'), ('year', 'i4'), ('amount', 'f8')])

# Map state letters to 0-9 indices
state_map = {'A':0, 'B':1, 'C':2, 'D':3, 'E':4, 'F':5, 'G':6, 'H':7, 'I':8, 'J':9}
state_indices = np.array([state_map[s] for s in customer_data['current_state']])

# Build initial state vectors in bulk: shape (n_customers, 10)
n_customers = len(customer_data)
initial_vectors = np.zeros((n_customers, 10), dtype=np.float64)
initial_vectors[np.arange(n_customers), state_indices] = customer_data['amount']

# Alternative one-liner for initial vectors (same result)
# initial_vectors = np.eye(10, dtype=np.float64)[state_indices] * customer_data['amount'][:, np.newaxis]

Step 2: Precompute Matrix Chains for Shared Start Years

Customers with the same start year use the exact same sequence of matrices. Instead of recalculating the matrix product sequence for each customer, we'll precompute these chains once per unique start year.

# Get start indices (customer's year → next matrix index, matches your original code)
start_indices = customer_data['year']

# Get all unique start years to avoid redundant work
unique_start_indices = np.unique(start_indices)

# Precompute matrix chains for each unique start index
prod_chain = {}
for t in unique_start_indices:
    # Slice the matrices from start index t to 13 (matches your original matrices[3:14])
    mats_slice = matrices[t:14]
    steps = len(mats_slice)
    
    # Initialize chain to store cumulative matrix products
    chain = np.zeros_like(mats_slice)
    chain[0] = mats_slice[0]
    
    # Build the cumulative product chain
    for i in range(1, steps):
        chain[i] = chain[i-1] @ mats_slice[i]
    
    prod_chain[t] = chain

Step 3: Batch Process Customers by Group

Now we'll process all customers in the same start year group at once using NumPy's broadcasted matrix multiplication—this is where the speed gain happens.

# Store results keyed by customer ID
all_results = {}

for t in unique_start_indices:
    # Filter customers with this start index
    mask = start_indices == t
    cust_ids = customer_data['customer'][mask]
    group_vectors = initial_vectors[mask]
    
    # Get the precomputed matrix chain for this start index
    chain = prod_chain[t]
    
    # Batch compute all steps for the group: shape (n_group_customers, steps, 10)
    # Broadcast the group vectors to match the chain's shape for matrix multiplication
    batch_results = group_vectors[:, np.newaxis, :] @ chain
    
    # Assign results to each customer
    for cust_id, results in zip(cust_ids, batch_results):
        all_results[cust_id] = list(results)

Why This Works Better Than Per-Customer Loops

  • Vectorized Operations: NumPy uses optimized BLAS/LAPACK libraries for matrix multiplication, which are way faster than Python-level loops over individual customers.
  • Redundant Calculation Elimination: By precomputing matrix chains for shared start years, we avoid recalculating the same matrix products thousands of times.

Validation Check

To make sure our batch results match the per-customer code you provided, let's spot-check one customer:

# Manual calculation for ax111 (matches your original code)
ax111_initial = np.array([300,0,0,0,0,0,0,0,0,0])
manual_results = []
output = ax111_initial
for arr in matrices[3:14]:
    output = output @ arr
    manual_results.append(output)

# Compare to batch result
batch_result = all_results['ax111']
print(np.allclose(manual_results, batch_result))  # Should print True

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 15:17:34