File size: 2,094 Bytes
be99550
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
#include <hip/hip_runtime.h>
#include <stdio.h>
#include <windows.h>

#define NUM_BRAINS 6

__global__ void moe_expert_kernel(uint32_t* vram, size_t size, int brain_id) {
    size_t idx = blockIdx.x * blockDim.x + threadIdx.x;
    if (idx < size) {
        uint32_t compressed_val = vram[idx];
        float sum = 0.0f;
        
        // ์‹ค์‹œ๊ฐ„ 16x ์••์ถ• ํ•ด์ œ (2-bit Unpacking) ๋ฐ 8-State ๋ณต์›
        #pragma unroll
        for(int j=0; j<16; j++) {
            uint32_t two_bits = (compressed_val >> (j * 2)) & 0x3;
            // ๋‡Œ(Expert)์˜ ์„ฑํ–ฅ(brain_id)์— ๋”ฐ๋ฅธ 8-State ํ™€๋กœ๊ทธ๋ž˜ํ”ฝ ์œ„์ƒ ๋งตํ•‘
            float decoded_weight = (float)two_bits - 1.5f + (brain_id * 0.1f);
            sum += decoded_weight;
        }
        
        // DCE(Dead Code Elimination) ๋ฐฉ์ง€๋ฅผ ์œ„ํ•ด ๋ณต์›๋œ ๊ฒฐ๊ณผ๊ฐ’์„ ๋‹ค์‹œ ๋ฉ”๋ชจ๋ฆฌ์— ์ €์žฅ
        vram[idx] = compressed_val ^ *((uint32_t*)&sum);
    }
}

int main() {
    size_t size = 10000000; 
    uint32_t* d_brains[NUM_BRAINS];
    
    for(int b=0; b<NUM_BRAINS; b++) hipMalloc(&d_brains[b], size * 4);
    
    LARGE_INTEGER freq, start, end;
    QueryPerformanceFrequency(&freq);
    
    for(int b=0; b<NUM_BRAINS; b++) {
        hipLaunchKernelGGL(moe_expert_kernel, dim3((size+255)/256), dim3(256), 0, 0, d_brains[b], size, b);
    }
    hipDeviceSynchronize();
    
    int passes = 1000; 
    
    QueryPerformanceCounter(&start);
    for(int i=0; i<passes; i++) {
        int target_brain = i % NUM_BRAINS; 
        hipLaunchKernelGGL(moe_expert_kernel, dim3((size+255)/256), dim3(256), 0, 0, d_brains[target_brain], size, target_brain);
    }
    hipDeviceSynchronize(); 
    
    uint32_t check_val;
    hipMemcpy(&check_val, d_brains[5], 4, hipMemcpyDeviceToHost);
    QueryPerformanceCounter(&end);
    
    double elapsed = (double)(end.QuadPart - start.QuadPart) / freq.QuadPart;
    double tps = 1000.0 / elapsed;
    
    printf(">> โฑ๏ธ True ROCm Time (w/ Decompression): %f s | ๐Ÿš€ Honest MoE TPS: %f\n", elapsed, tps);
    for(int b=0; b<NUM_BRAINS; b++) hipFree(d_brains[b]);
    return 0;
}