Multi-Agent Simulation
2026/6/6大约 1 分钟
Multi-Agent Simulation
题目描述
实现多智能体集群模拟(Boids 算法)。输入包含 个智能体,每个智能体由 4 个连续 32 位浮点数表示 (位置和速度)。总数组大小为 。
Boids 模型包含三种行为规则:分离(避开邻近智能体)、对齐(匹配邻近智能体的平均速度方向)、凝聚(向邻近智能体的质心移动)。模拟每个时间步更新所有智能体的位置和速度。
实现要求
- 不允许使用外部库。
solve函数签名必须保持不变。
约束条件
- ,。
解题思路
每个时间步需要计算所有智能体对之间的相互作用( 朴素复杂度)。GPU 上可利用空间哈希网格(Spatial Hashing)将搜索限制在邻近单元内,降低复杂度。每个智能体由一个线程处理:先查询邻近智能体,再计算三种力并更新速度和位置。属于典型的 N-body 问题变体。
代码实现
CUDA
#include <cuda_runtime.h>
__global__ void boids_step(float* agents, int N, float dt, float sep_r, float ali_r, float coh_r) {
int i=blockIdx.x*blockDim.x+threadIdx.x; if(i>=N)return;
float xi=agents[i*4],yi=agents[i*4+1],vxi=agents[i*4+2],vyi=agents[i*4+3];
float sx=0,sy=0,svx=0,svy=0,scx=0,scy=0; int ns=0,na=0,nc=0;
for(int j=0;j<N;j++)if(j!=i){
float dx=xi-agents[j*4],dy=yi-agents[j*4+1]; float d2=dx*dx+dy*dy;
if(d2<sep_r*sep_r){sx+=dx;sy+=dy;ns++;} // Separation
if(d2<ali_r*ali_r){svx+=agents[j*4+2];svy+=agents[j*4+3];na++;} // Alignment
if(d2<coh_r*coh_r){scx+=agents[j*4];scy+=agents[j*4+1];nc++;} // Cohesion
}
float fx=0,fy=0;
if(ns>0){fx+=sx/ns;fy+=sy/ns;} // Separation force
if(na>0){fx+=(svx/na-vxi)*0.5f;fy+=(svy/na-vyi)*0.5f;} // Alignment
if(nc>0){fx+=(scx/nc-xi)*0.01f;fy+=(scy/nc-yi)*0.01f;} // Cohesion
agents[i*4+2]=vxi+fx*dt; agents[i*4+3]=vyi+fy*dt;
agents[i*4]=xi+agents[i*4+2]*dt; agents[i*4+1]=yi+agents[i*4+3]*dt;
}
extern "C" void solve(float* agents, int N, float dt, float sep_r, float ali_r, float coh_r) {
boids_step<<<(N+255)/256,256>>>(agents,N,dt,sep_r,ali_r,coh_r);
cudaDeviceSynchronize();
}Triton
import triton, triton.language as tl
@triton.jit
def boids_step(agents_ptr, N,dt,sep_r,ali_r,coh_r, BLOCK:tl.constexpr):
i=tl.program_id(0)*BLOCK+tl.arange(0,BLOCK); mask=i<N
xi=tl.load(agents_ptr+i*4,mask=mask);yi=tl.load(agents_ptr+i*4+1,mask=mask)
vxi=tl.load(agents_ptr+i*4+2,mask=mask);vyi=tl.load(agents_ptr+i*4+3,mask=mask)
fx=tl.zeros((BLOCK,),tl.float32);fy=tl.zeros((BLOCK,),tl.float32)
for j in range(N):
xj=tl.load(agents_ptr+j*4);yj=tl.load(agents_ptr+j*4+1)
dx=xi-xj;dy=yi-yj;d2=dx*dx+dy*dy; me=i!=j
fx+=tl.where(d2<sep_r*sep_r&me,dx,0.0); fy+=tl.where(d2<sep_r*sep_r&me,dy,0.0)
vxi+=fx*dt;vyi+=fy*dt; xi+=vxi*dt;yi+=vyi*dt
tl.store(agents_ptr+i*4,xi,mask=mask);tl.store(agents_ptr+i*4+1,yi,mask=mask)
tl.store(agents_ptr+i*4+2,vxi,mask=mask);tl.store(agents_ptr+i*4+3,vyi,mask=mask)