Pull requests / #1737

#1737 route-resident: one warp per token in the swap kernel

open · @aly8246 · 0 コメント · GitHub で見る

Models & quants

本文

route-resident: one warp per token in the swap kernel

The swap kernel is the only work STRATA_ROUTE_RESIDENT adds to a window, and it was written as one
thread per token: with a verify window of 1 to 8 rows that is up to 8 threads of a block, each doing,
per rank in [lo, hi], a 512-expert scan with a ten-way "already picked" test inside it.  On this
machine that is about 4 x 512 x 12 scalar operations running on six threads, on the critical stream,
once per layer, and it shows up as its own cost in the window.

STRATA_VERIFY_PROFILE, three legs in one session, same prompts, same expert cache (8083 slots, one
learned profile), medians of the per-window stage table:

    configuration                        hc-read1+router   waitB    ms/window   pcie_e
    option off                                2.25          9.30      21.85      1.87
    option on, one thread per token           5.00          5.01      21.97      1.00
    option on, this commit                    2.58          4.87      19.75      0.95

The option redirects about a miss per layer-window away from PCIe and saves 4.3 ms of waiting for it;
the serial swap spends 2.75 ms back, which is why the option measured as a wash.  In the warp form the
swap costs 0.33 ms, so the waiting it saved stays saved.

End to end, twelve matched pairs of engine restarts, nine decode samples a leg, three greedy 150-token
prompts, the learned expert profile copied back before every leg so both arms start from the same
8083-slot cache.  The first six pairs ran the unpatched build first, the last six the patched build
first, so a within-pair order effect would show up as the two halves disagreeing:

    pair  order   decode unpatched   patched    delta    ms/window unpatched   patched    delta
     1    A->B         104.1          115.7     +11.1%         18.64          17.01      -8.7%
     2    A->B         103.5          107.4      +3.8%         18.47          18.12      -1.9%
     3    A->B         111.4          107.7      -3.3%         17.71          18.26      +3.1%
     4    A->B         101.4          113.0     +11.4%         19.11          17.03     -10.9%
     5    A->B         105.8          107.6      +1.7%         19.05          18.05      -5.2%
     6    A->B         104.8          105.6      +0.8%         18.60          17.83      -4.1%
     7    B->A         106.6          113.4      +6.4%         18.70          17.16      -8.2%
     8    B->A         103.1          110.2      +6.9%         18.14          17.29      -4.7%
     9    B->A         106.9          111.5      +4.3%         17.91          17.26      -3.6%
    10    B->A         104.3          114.7     +10.0%         17.81          17.11      -3.9%
    11    B->A         108.3          115.9      +7.0%         18.47          16.61     -10.1%
    12    B->A         104.5          113.6      +8.7%         18.43          17.76      -3.6%

    twelve legs each:   unpatched  decode median 104.7 (101.4-111.4)   ms/window median 18.47 (17.71-19.11)
                        this commit decode median 112.2 (105.6-115.9)  ms/window median 17.27 (16.61-18.26)

    decode      median +6.6%   mean +5.7%   faster in 11 of 12 pairs   (sign test p = 0.006)
    ms/window   median -4.4%   mean -5.2%   shorter in 11 of 12 pairs  (sign test p = 0.006)
    first six pairs: +2.7% decode, -4.7% window     last six: +7.0% decode, -4.3% window

pcie_e, the engine's own count of experts per layer-window that cross PCIe, is the same in both arms
(leg medians 0.63 against 0.64), so the same experts are routed and the difference is the kernel.

The swap's cost, and so its removal, scales with the tail experts the option actually redirects.  Four
more pairs with a 4000-slot cache instead of the automatic 8083 (pcie_e 1.83 against 1.78, i.e. about
three times the traffic of the series above), two pairs each way:

    pair  order   decode unpatched   patched    delta    ms/window unpatched   patched    delta
     1    X->Y          83.8           92.0      +9.8%         21.94          20.60      -6.1%
     2    Y->X          84.6           92.5      +9.3%         23.97          20.53     -14.4%
     3    X->Y          82.2           85.5      +4.0%         23.69          22.65      -4.4%
     4    Y->X          81.0           87.7      +8.3%         24.54          22.21      -9.5%

    decode median +8.8%, ms/window median -7.8%, 4 of 4 pairs, both orders.

What this changes is the shape, not the choice.  One warp per token: lane j reads experts j, j+32, ...
(coalesced instead of a serial walk), the warp reduces to the best logit in a shuffle tree, and the
ranks are still visited in order, because a swap changes the set already picked.  The choice is the
same bits: a strictly greater logit wins, the smallest expert index breaks ties, the already-picked
test runs against the ids as they are at that rank, a token with no swap keeps the router's weights
untouched, and the softmax is the same expression.

Parity.  400 random cases - half of them built with many equal logits, random lo/hi, margins from
-0.5 to 1.5, 1 to 8 tokens, about 60% of the experts resident, routed ids with repeats - compare the
two implementations bit for bit on the ids, on the weights (their raw bits) and on the four counters.
The harness carries both kernels side by side; 0 mismatches.

For context, the option's price is the margin's and not the kernel's, and this commit does not touch
it.  Teacher-forced log-probabilities over 7677 positions of one 7.6k-token text (technical prose,
source code and Chinese), top-64, against the same text with the option off:

    margin   argmax same   KL (top-64)   teacher-forced loss   perplexity
    off      -             0.000         3.2033                -
    0.25     84.9%         0.235         3.2238                +0.6%
    0.50     83.0%         0.293         3.4360                +7.3%
    1.00     63.5%         1.001         4.7751                +49%

Two runs with the option off are bit-identical, so that spread is the option and not the machine.
0.25 keeps nearly all of the waitB saving for a fraction of the distortion; the default 1.0 is not
worth its price.

Tests: the parity harness is a standalone .cu (both kernels, the random cases, the bit comparison).
Happy to wire it into the test suite next to the other kernel tests if you want it in the tree.

関連リンク

インストール・モデル・リリースへの站内リンク。