chore(repo): 瘦身规范 - 忽略编译产物/JIT缓存/raw_outputs,取消跟踪618个垃圾文件
- .gitignore 新增规范:实验目录只提交 代码 + report.md + results.json(小体量汇总) - 忽略: *.o/*.so/*.ninja*/*.cubin, sglang_sm120_cache/, vllm_sm120_cache/, *_sm120_cache/, experiments/**/runtime/, bench-output/, experiments/**/raw_outputs/ - 修正原 experiments/*/runtime/ 规则过窄(只匹配一层) -> experiments/**/runtime/ - git rm --cached 取消跟踪 618 个已入库的缓存/产物/原始日志(文件保留在本地工作区) - 当前 HEAD 减少约 83MB 跟踪体积; 历史体积需另做 filter-branch 重写(本次不做) - results.json 保留(28个共0.99MB, compare.py/文档依赖) Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
parent
7231536db2
commit
080e095ece
27
.gitignore
vendored
27
.gitignore
vendored
@ -42,7 +42,7 @@ experiments/*/results/*/*.tsv
|
||||
experiments/*/results/*/*/*.tsv
|
||||
|
||||
# 运行时目录(pid 文件、缓存、临时文件)
|
||||
experiments/*/runtime/
|
||||
experiments/**/runtime/
|
||||
|
||||
# 无关项目
|
||||
loomeval_yy/
|
||||
@ -52,3 +52,28 @@ loomeval_yy/
|
||||
dsv4_dspark_h20_sglang_tp_dp_matrix
|
||||
dsv4_h20_sglang_tp_dp_matrix
|
||||
glm_old_h20_vllm_tp_dp_matrix
|
||||
# === 仓库瘦身规范(2026-07)===
|
||||
# 约定:实验目录只提交 代码(config.env/*.sh/*.py) + report.md + results.json(小体量汇总)。
|
||||
# 不提交:逐请求原始日志(raw_outputs)、JIT/算子缓存、编译产物、一次性 bench 输出。
|
||||
# 历史中 raw_outputs/*.jsonl 是仓库膨胀主因(~1.3GB);本规范从源头不再入库。
|
||||
# results.json 须保持小体量,禁止嵌入 raw_requests(否则单独走 raw_outputs)。
|
||||
|
||||
# 编译产物(每次构建重新生成)
|
||||
*.o
|
||||
*.so
|
||||
*.ninja
|
||||
*.ninja_deps
|
||||
*.ninja_log
|
||||
*.cubin
|
||||
build/
|
||||
|
||||
# JIT / 算子缓存(每次跑重新生成)
|
||||
sglang_sm120_cache/
|
||||
vllm_sm120_cache/
|
||||
*_sm120_cache/
|
||||
|
||||
# 一次性 bench 输出
|
||||
bench-output/
|
||||
|
||||
# 逐请求原始日志(体积大;汇总见 results.json / report.md)
|
||||
experiments/**/raw_outputs/
|
||||
|
||||
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 64; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 64; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 11; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 11; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 19; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 19; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 6; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 6; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 15; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 15; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,74 +0,0 @@
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_post_tilelang_kernel(const float* a, const bfloat16_t* b, const float* c, const bfloat16_t* d, bfloat16_t* x, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(128, 1) mhc_post_tilelang_kernel(const float* a, const bfloat16_t* b, const float* c, const bfloat16_t* d, bfloat16_t* x, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float a_local[16];
|
||||
float c_local[4];
|
||||
float b_local[32];
|
||||
bfloat16_t b_shared_local_cast[8];
|
||||
bfloat16_t d_shared_local_cast_1[8];
|
||||
float d_local[8];
|
||||
float x_local[32];
|
||||
bfloat16_t x_local_cast_2[8];
|
||||
cudaGridDependencySynchronize();
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
*(float4*)(a_local + (i * 4)) = *(float4*)(a + ((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)i) * (int64_t)4)));
|
||||
}
|
||||
*(float4*)(c_local + 0) = *(float4*)(c + (((int64_t)((int)blockIdx.x)) * (int64_t)4));
|
||||
for (int i0_h = 0; i0_h < 4; ++i0_h) {
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 4; ++i_1) {
|
||||
*(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((i_1 * 1024) + (((int)threadIdx.x) * 8))) = *(uint4*)(b + ((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)i_1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)));
|
||||
}
|
||||
*(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((int)threadIdx.x) * 8) + 4096)) = *(uint4*)(d + (((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h) * (int64_t)1024)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)));
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 4; ++i_2) {
|
||||
*(uint4*)(b_shared_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((i_2 * 1024) + (((int)threadIdx.x) * 8)));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __1;
|
||||
uint2 v_ = *(uint2*)(b_shared_local_cast + (vec * 4));
|
||||
((float2*)(&__1))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__1))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(b_local + ((i_2 * 8) + (vec * 4))) = __1;
|
||||
}
|
||||
}
|
||||
*(uint4*)(d_shared_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((int)threadIdx.x) * 8) + 4096));
|
||||
for (int i_3 = 0; i_3 < 2; ++i_3) {
|
||||
float4 __2;
|
||||
uint2 v__1 = *(uint2*)(d_shared_local_cast_1 + (i_3 * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__1))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__1))[1]);
|
||||
*(float4*)(d_local + (i_3 * 4)) = __2;
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 32; ++i_4) {
|
||||
x_local[i_4] = (c_local[(i_4 >> 3)] * d_local[(i_4 & 7)]);
|
||||
for (int i_hci = 0; i_hci < 4; ++i_hci) {
|
||||
x_local[i_4] = (x_local[i_4] + (a_local[((i_hci * 4) + (i_4 >> 3))] * b_local[((i_hci * 8) + (i_4 & 7))]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
uint2 __3;
|
||||
float4 v__2 = *(float4*)(x_local + ((i_5 * 8) + (vec_1 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__3))[0] = __float22bfloat162_rn(((float2*)(&v__2))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__3))[1] = __float22bfloat162_rn(((float2*)(&v__2))[1]);
|
||||
*(uint2*)(x_local_cast_2 + (vec_1 * 4)) = __3;
|
||||
}
|
||||
*(uint4*)(x + ((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)i_5) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8))) = *(uint4*)(x_local_cast_2 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 26; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 26; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 13; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 13; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 39; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 39; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 8; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 8; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 4; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 4; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 2; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 2; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,420 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((int64_t)((int)blockIdx.x))]);
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[((((int64_t)((int)blockIdx.x)) * (int64_t)24) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 5; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 5; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,104 +0,0 @@
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_fused_tilelang_kernel(const float* comb_mix, const float* post_mix, const bfloat16_t* residual_in, bfloat16_t* residual_out, float* rp_out, const float* weight_t, const bfloat16_t* x_in, float* yp_out, int num_tokens, int split_k);
|
||||
extern "C" __global__ void __launch_bounds__(256, 1) mhc_fused_tilelang_kernel(const float* comb_mix, const float* post_mix, const bfloat16_t* residual_in, bfloat16_t* residual_out, float* rp_out, const float* weight_t, const bfloat16_t* x_in, float* yp_out, int num_tokens, int split_k) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float acc[3];
|
||||
float sqr[1];
|
||||
float pm[4];
|
||||
float cm[16];
|
||||
float new_r[4];
|
||||
float v = 0x0p+0f/*0.000000e+00*/;
|
||||
float v2 = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 3; ++i) {
|
||||
acc[i] = 0x0p+0f/*0.000000e+00*/;
|
||||
}
|
||||
sqr[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
((float*)buf_dyn_shmem)[((int)threadIdx.x)] = post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))];
|
||||
}
|
||||
if (((int)threadIdx.x) < 16) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] = comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + ((int64_t)((int)threadIdx.x)))];
|
||||
}
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 4; ++j) {
|
||||
pm[j] = ((float*)buf_dyn_shmem)[j];
|
||||
}
|
||||
#pragma unroll
|
||||
for (int j_1 = 0; j_1 < 4; ++j_1) {
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 4; ++k) {
|
||||
cm[((k * 4) + j_1)] = ((float*)buf_dyn_shmem)[(((k * 4) + j_1) + 4)];
|
||||
}
|
||||
}
|
||||
for (int it = 0; it < ((4096 / split_k) >> 8); ++it) {
|
||||
#pragma unroll
|
||||
for (int j_2 = 0; j_2 < 4; ++j_2) {
|
||||
new_r[j_2] = (pm[j_2] * ((float)x_in[((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))]));
|
||||
#pragma unroll
|
||||
for (int k_1 = 0; k_1 < 4; ++k_1) {
|
||||
new_r[j_2] = (new_r[j_2] + (cm[((k_1 * 4) + j_2)] * ((float)residual_in[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)k_1) * (int64_t)4096)) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))])));
|
||||
}
|
||||
}
|
||||
if (((int)blockIdx.y) == 0) {
|
||||
#pragma unroll
|
||||
for (int j_3 = 0; j_3 < 4; ++j_3) {
|
||||
residual_out[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)j_3) * (int64_t)4096)) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))] = ((bfloat16_t)new_r[j_3]);
|
||||
sqr[0] = (sqr[0] + (new_r[j_3] * new_r[j_3]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int n = 0; n < 3; ++n) {
|
||||
#pragma unroll
|
||||
for (int j_4 = 0; j_4 < 4; ++j_4) {
|
||||
acc[n] = (acc[n] + (weight_t[((((((((int64_t)((int)blockIdx.y)) * (int64_t)49152) + (((int64_t)n) * (int64_t)16384)) + (((int64_t)j_4) * (int64_t)4096)) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))] * new_r[j_4]));
|
||||
}
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int n_1 = 0; n_1 < 3; ++n_1) {
|
||||
acc[n_1] = tl::warp_reduce_sum(acc[n_1]);
|
||||
}
|
||||
if (((int)blockIdx.y) == 0) {
|
||||
sqr[0] = tl::warp_reduce_sum(sqr[0]);
|
||||
}
|
||||
if ((((int)threadIdx.x) % 32) == 0) {
|
||||
#pragma unroll
|
||||
for (int n_2 = 0; n_2 < 3; ++n_2) {
|
||||
((float*)buf_dyn_shmem)[((((((int)threadIdx.x) >> 5) * 4) + n_2) + 20)] = acc[n_2];
|
||||
}
|
||||
if (((int)blockIdx.y) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((((int)threadIdx.x) >> 5) * 4) + 23)] = sqr[0];
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
if ((((int)threadIdx.x) >> 5) == 0) {
|
||||
if ((((int)threadIdx.x) & 31) < 3) {
|
||||
#pragma unroll
|
||||
for (int w = 0; w < 8; ++w) {
|
||||
v = (v + ((float*)buf_dyn_shmem)[(((w * 4) + (((int)threadIdx.x) & 31)) + 20)]);
|
||||
}
|
||||
yp_out[((((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)((int)blockIdx.z)) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)blockIdx.y)) * (int64_t)3)) + (((int64_t)((int)threadIdx.x)) & (int64_t)31))] = v;
|
||||
}
|
||||
if ((((int)blockIdx.y) == 0) && ((((int)threadIdx.x) % 32) == 0)) {
|
||||
#pragma unroll
|
||||
for (int w_1 = 0; w_1 < 8; ++w_1) {
|
||||
v2 = (v2 + ((float*)buf_dyn_shmem)[((w_1 * 4) + 23)]);
|
||||
}
|
||||
rp_out[((((int64_t)((int)blockIdx.z)) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))] = v2;
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,102 +0,0 @@
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_fused_tilelang_kernel(const float* comb_mix, const float* post_mix, const bfloat16_t* residual_in, bfloat16_t* residual_out, float* rp_out, const float* weight_t, const bfloat16_t* x_in, float* yp_out, int num_tokens, int split_k);
|
||||
extern "C" __global__ void __launch_bounds__(256, 1) mhc_fused_tilelang_kernel(const float* comb_mix, const float* post_mix, const bfloat16_t* residual_in, bfloat16_t* residual_out, float* rp_out, const float* weight_t, const bfloat16_t* x_in, float* yp_out, int num_tokens, int split_k) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float acc[2];
|
||||
float sqr[1];
|
||||
float pm[4];
|
||||
float cm[16];
|
||||
float new_r[4];
|
||||
float v = 0x0p+0f/*0.000000e+00*/;
|
||||
float v2 = 0x0p+0f/*0.000000e+00*/;
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float2*)(acc + 0) = make_float2(broadcast_var, broadcast_var);
|
||||
sqr[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
((float*)buf_dyn_shmem)[((int)threadIdx.x)] = post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))];
|
||||
}
|
||||
if (((int)threadIdx.x) < 16) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] = comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + ((int64_t)((int)threadIdx.x)))];
|
||||
}
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int j = 0; j < 4; ++j) {
|
||||
pm[j] = ((float*)buf_dyn_shmem)[j];
|
||||
}
|
||||
#pragma unroll
|
||||
for (int j_1 = 0; j_1 < 4; ++j_1) {
|
||||
#pragma unroll
|
||||
for (int k = 0; k < 4; ++k) {
|
||||
cm[((k * 4) + j_1)] = ((float*)buf_dyn_shmem)[(((k * 4) + j_1) + 4)];
|
||||
}
|
||||
}
|
||||
for (int it = 0; it < ((4096 / split_k) >> 8); ++it) {
|
||||
#pragma unroll
|
||||
for (int j_2 = 0; j_2 < 4; ++j_2) {
|
||||
new_r[j_2] = (pm[j_2] * ((float)x_in[((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))]));
|
||||
#pragma unroll
|
||||
for (int k_1 = 0; k_1 < 4; ++k_1) {
|
||||
new_r[j_2] = (new_r[j_2] + (cm[((k_1 * 4) + j_2)] * ((float)residual_in[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)k_1) * (int64_t)4096)) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))])));
|
||||
}
|
||||
}
|
||||
if (((int)blockIdx.y) == 0) {
|
||||
#pragma unroll
|
||||
for (int j_3 = 0; j_3 < 4; ++j_3) {
|
||||
residual_out[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)j_3) * (int64_t)4096)) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))] = ((bfloat16_t)new_r[j_3]);
|
||||
sqr[0] = (sqr[0] + (new_r[j_3] * new_r[j_3]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int n = 0; n < 2; ++n) {
|
||||
#pragma unroll
|
||||
for (int j_4 = 0; j_4 < 4; ++j_4) {
|
||||
acc[n] = (acc[n] + (weight_t[((((((((int64_t)((int)blockIdx.y)) * (int64_t)32768) + (((int64_t)n) * (int64_t)16384)) + (((int64_t)j_4) * (int64_t)4096)) + (((int64_t)it) * (int64_t)256)) + (((int64_t)((int)blockIdx.z)) * ((int64_t)4096 / ((int64_t)split_k)))) + ((int64_t)((int)threadIdx.x)))] * new_r[j_4]));
|
||||
}
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int n_1 = 0; n_1 < 2; ++n_1) {
|
||||
acc[n_1] = tl::warp_reduce_sum(acc[n_1]);
|
||||
}
|
||||
if (((int)blockIdx.y) == 0) {
|
||||
sqr[0] = tl::warp_reduce_sum(sqr[0]);
|
||||
}
|
||||
if ((((int)threadIdx.x) % 32) == 0) {
|
||||
#pragma unroll
|
||||
for (int n_2 = 0; n_2 < 2; ++n_2) {
|
||||
((float*)buf_dyn_shmem)[((((((int)threadIdx.x) >> 5) * 3) + n_2) + 20)] = acc[n_2];
|
||||
}
|
||||
if (((int)blockIdx.y) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((((int)threadIdx.x) >> 5) * 3) + 22)] = sqr[0];
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
if ((((int)threadIdx.x) >> 5) == 0) {
|
||||
if ((((int)threadIdx.x) & 31) < 2) {
|
||||
#pragma unroll
|
||||
for (int w = 0; w < 8; ++w) {
|
||||
v = (v + ((float*)buf_dyn_shmem)[(((w * 3) + (((int)threadIdx.x) & 31)) + 20)]);
|
||||
}
|
||||
yp_out[((((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)((int)blockIdx.z)) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)blockIdx.y)) * (int64_t)2)) + (((int64_t)((int)threadIdx.x)) & (int64_t)31))] = v;
|
||||
}
|
||||
if ((((int)blockIdx.y) == 0) && ((((int)threadIdx.x) % 32) == 0)) {
|
||||
#pragma unroll
|
||||
for (int w_1 = 0; w_1 < 8; ++w_1) {
|
||||
v2 = (v2 + ((float*)buf_dyn_shmem)[((w_1 * 3) + 22)]);
|
||||
}
|
||||
rp_out[((((int64_t)((int)blockIdx.z)) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))] = v2;
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,190 +0,0 @@
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void hc_head_fuse_tilelang_kernel(const float* fn, const float* hc_base, const float* hc_scale, bfloat16_t* out, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(128, 1) hc_head_fuse_tilelang_kernel(const float* fn, const float* hc_base, const float* hc_scale, bfloat16_t* out, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float sqrsum_r[1];
|
||||
float mixes_r[4];
|
||||
bfloat16_t residual_local_cast[8];
|
||||
float x_local[8];
|
||||
float fn_local[8];
|
||||
float rsqrt_val[1];
|
||||
float xl[32];
|
||||
float ol[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t out_local_cast_2[8];
|
||||
bfloat16_t xs_local_cast_3[8];
|
||||
bfloat16_t out_local_cast_4[8];
|
||||
bfloat16_t xs_local_cast_5[8];
|
||||
bfloat16_t out_local_cast_6[8];
|
||||
cudaGridDependencySynchronize();
|
||||
sqrsum_r[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(mixes_r + 0) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
for (int m_c = 0; m_c < 4; ++m_c) {
|
||||
for (int i_h = 0; i_h < 4; ++i_h) {
|
||||
*(uint4*)(residual_local_cast + 0) = *(uint4*)(residual + ((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)m_c) * (int64_t)4096)) + (((int64_t)i_h) * (int64_t)1024)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)));
|
||||
for (int i = 0; i < 2; ++i) {
|
||||
float4 __1;
|
||||
uint2 v_ = *(uint2*)(residual_local_cast + (i * 4));
|
||||
((float2*)(&__1))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__1))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(x_local + (i * 4)) = __1;
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
sqrsum_r[0] = (sqrsum_r[0] + (x_local[i_1] * x_local[i_1]));
|
||||
}
|
||||
#pragma unroll
|
||||
for (int m_m = 0; m_m < 4; ++m_m) {
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 2; ++i_2) {
|
||||
*(float4*)(fn_local + (i_2 * 4)) = *(float4*)(fn + (((((m_m * 16384) + (m_c * 4096)) + (i_h * 1024)) + (((int)threadIdx.x) * 8)) + (i_2 * 4)));
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
mixes_r[m_m] = (mixes_r[m_m] + (x_local[i_3] * fn_local[i_3]));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
sqrsum_r[0] = tl::AllReduce<tl::SumOp, 128, 1, 0, tl::NamedBarrier<128>>::run(sqrsum_r[0], (&(((float*)buf_dyn_shmem)[0])));
|
||||
__syncthreads();
|
||||
for (int __finred_0 = 0; __finred_0 < 4; ++__finred_0) {
|
||||
mixes_r[__finred_0] = tl::AllReduce<tl::SumOp, 128, 1, 0, tl::NamedBarrier<128>>::run(mixes_r[__finred_0], (&(((float*)buf_dyn_shmem)[0])));
|
||||
}
|
||||
rsqrt_val[0] = rsqrtf(((sqrsum_r[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
__syncthreads();
|
||||
if ((((int)threadIdx.x) >> 2) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) & 3)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - (((mixes_r[(((int)threadIdx.x) & 3)] * rsqrt_val[0]) * hc_scale[0]) + hc_base[(((int)threadIdx.x) & 3)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 4; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_4 * 1024) + (((int)threadIdx.x) * 8)) + 256)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)i_4) * (int64_t)4096)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_5 * 1024) + (((int)threadIdx.x) * 8)) + 4352)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)i_5) * (int64_t)4096)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 4; ++i_6) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_6 * 1024)) + (((int)threadIdx.x) * 8)) + 256));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v__1 = *(uint2*)(xs_local_cast_1 + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__1))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__1))[1]);
|
||||
*(float4*)(xl + ((i_6 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + (i_7 * 1024)) + (((int)threadIdx.x) * 8)) + 256)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + (((int64_t)i_7) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 2; ++i_8) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_8 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
__syncthreads();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[i_hc];
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 8; ++i_9) {
|
||||
ol[i_9] = (ol[i_9] + (pre * xl[((i_hc * 8) + i_9)]));
|
||||
}
|
||||
}
|
||||
for (int i_10 = 0; i_10 < 2; ++i_10) {
|
||||
uint2 __3;
|
||||
float4 v__2 = *(float4*)(ol + (i_10 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__3))[0] = __float22bfloat162_rn(((float2*)(&v__2))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__3))[1] = __float22bfloat162_rn(((float2*)(&v__2))[1]);
|
||||
*(uint2*)(out_local_cast_2 + (i_10 * 4)) = __3;
|
||||
}
|
||||
*(uint4*)(out + (((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h) * (int64_t)1024)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8))) = *(uint4*)(out_local_cast_2 + 0);
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
*(uint4*)(xs_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_11 * 1024) + (((int)threadIdx.x) * 8)) + 256));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __4;
|
||||
uint2 v__3 = *(uint2*)(xs_local_cast_3 + (vec_1 * 4));
|
||||
((float2*)(&__4))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__3))[0]);
|
||||
((float2*)(&__4))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__3))[1]);
|
||||
*(float4*)(xl + ((i_11 * 8) + (vec_1 * 4))) = __4;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 2; ++i_12) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_12 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[i_hc_1];
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 8; ++i_13) {
|
||||
ol[i_13] = (ol[i_13] + (pre_1 * xl[((i_hc_1 * 8) + i_13)]));
|
||||
}
|
||||
}
|
||||
for (int i_14 = 0; i_14 < 2; ++i_14) {
|
||||
uint2 __5;
|
||||
float4 v__4 = *(float4*)(ol + (i_14 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__4))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__4))[1]);
|
||||
*(uint2*)(out_local_cast_4 + (i_14 * 4)) = __5;
|
||||
}
|
||||
*(uint4*)(out + (((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2048)) = *(uint4*)(out_local_cast_4 + 0);
|
||||
tl::cp_async_wait<0>();
|
||||
__syncthreads();
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
*(uint4*)(xs_local_cast_5 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_15 * 1024) + (((int)threadIdx.x) * 8)) + 4352));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __6;
|
||||
uint2 v__5 = *(uint2*)(xs_local_cast_5 + (vec_2 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__5))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__5))[1]);
|
||||
*(float4*)(xl + ((i_15 * 8) + (vec_2 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_16 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[i_hc_2];
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 8; ++i_17) {
|
||||
ol[i_17] = (ol[i_17] + (pre_2 * xl[((i_hc_2 * 8) + i_17)]));
|
||||
}
|
||||
}
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
uint2 __7;
|
||||
float4 v__6 = *(float4*)(ol + (i_18 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__7))[0] = __float22bfloat162_rn(((float2*)(&v__6))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__7))[1] = __float22bfloat162_rn(((float2*)(&v__6))[1]);
|
||||
*(uint2*)(out_local_cast_6 + (i_18 * 4)) = __7;
|
||||
}
|
||||
*(uint4*)(out + (((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)3072)) = *(uint4*)(out_local_cast_6 + 0);
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 9; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 9; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 3; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 3; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Binary file not shown.
@ -1,424 +0,0 @@
|
||||
#include <math_constants.h>
|
||||
#include <tl_templates/cuda/gemm.h>
|
||||
#include <tl_templates/cuda/copy.h>
|
||||
#include <tl_templates/cuda/reduce.h>
|
||||
#include <tl_templates/cuda/ldsm.h>
|
||||
#include <tl_templates/cuda/threadblock_swizzle.h>
|
||||
#include <tl_templates/cuda/debug.h>
|
||||
#ifdef ENABLE_BF16
|
||||
#include <tl_templates/cuda/cuda_bf16_fallbacks.cuh>
|
||||
#endif
|
||||
|
||||
extern "C" __global__ void mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens);
|
||||
extern "C" __global__ void __launch_bounds__(96, 1) mhc_pre_big_fuse_with_norm_tilelang_kernel(float* comb_mix, const float* gemm_out_mul, const float* gemm_out_sqrsum, const float* hc_base, const float* hc_scale, bfloat16_t* layer_input, const bfloat16_t* norm_weight, float* post_mix, const bfloat16_t* residual, int num_tokens) {
|
||||
extern __shared__ __align__(1024) uchar buf_dyn_shmem[];
|
||||
float mixes[1];
|
||||
float rms[1];
|
||||
float cm[1];
|
||||
float row_max[1];
|
||||
float row_sum[1];
|
||||
float col_sum[1];
|
||||
float sumsq_per_pos[16];
|
||||
float sumsq[1];
|
||||
float rsqrt_norm[1];
|
||||
float xl[64];
|
||||
float ol[16];
|
||||
bfloat16_t xs_local_cast[8];
|
||||
bfloat16_t xs_local_cast_1[8];
|
||||
bfloat16_t xs_local_cast_2[8];
|
||||
float w_local[16];
|
||||
float ol_1[16];
|
||||
bfloat16_t w_shared_local_cast_3[8];
|
||||
bfloat16_t output_shared_local_cast_4[8];
|
||||
bfloat16_t layer_input_local_cast_5[8];
|
||||
bfloat16_t w_shared_local_cast_6[8];
|
||||
bfloat16_t output_shared_local_cast_7[8];
|
||||
bfloat16_t layer_input_local_cast_8[8];
|
||||
bfloat16_t w_shared_local_cast_9[8];
|
||||
bfloat16_t output_shared_local_cast_10[8];
|
||||
bfloat16_t layer_input_local_cast_11[8];
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
rms[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
cudaGridDependencySynchronize();
|
||||
for (int i_split = 0; i_split < 7; ++i_split) {
|
||||
rms[0] = (rms[0] + gemm_out_sqrsum[((((int64_t)i_split) * ((int64_t)num_tokens)) + ((int64_t)((int)blockIdx.x)))]);
|
||||
}
|
||||
rms[0] = rsqrtf(((rms[0] / 0x1p+14f/*1.638400e+04*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
mixes[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
for (int i_split_1 = 0; i_split_1 < 7; ++i_split_1) {
|
||||
mixes[0] = (mixes[0] + gemm_out_mul[(((((int64_t)((int)blockIdx.x)) * (int64_t)24) + ((((int64_t)i_split_1) * ((int64_t)num_tokens)) * (int64_t)24)) + (((int64_t)((int)threadIdx.x)) % (int64_t)24))]);
|
||||
}
|
||||
mixes[0] = (mixes[0] * rms[0]);
|
||||
if ((((int)threadIdx.x) / 24) == 0) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) % 24)] = mixes[0];
|
||||
}
|
||||
__syncthreads();
|
||||
if (((int)threadIdx.x) < 32) {
|
||||
if (((int)threadIdx.x) < 4) {
|
||||
post_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)4) + ((int64_t)((int)threadIdx.x)))] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 4)] * hc_scale[1]) + hc_base[(((int)threadIdx.x) + 4)]))))) * 0x1p+1f/*2.000000e+00*/);
|
||||
}
|
||||
cm[0] = ((((float*)buf_dyn_shmem)[((((int)threadIdx.x) & 15) + 8)] * hc_scale[2]) + hc_base[((((int)threadIdx.x) & 15) + 8)]);
|
||||
row_max[0] = -CUDART_INF_F;
|
||||
row_max[0] = max(row_max[0], cm[0]);
|
||||
row_max[0] = tl::AllReduce<tl::MaxOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_max[0]);
|
||||
cm[0] = expf((cm[0] - row_max[0]));
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = ((cm[0] / row_sum[0]) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
for (int __1 = 0; __1 < 19; ++__1) {
|
||||
row_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
row_sum[0] = (row_sum[0] + cm[0]);
|
||||
row_sum[0] = tl::AllReduce<tl::SumOp, 4, 1, 0, tl::NamedBarrier<32>>::run(row_sum[0]);
|
||||
cm[0] = (cm[0] / (row_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
col_sum[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
col_sum[0] = (col_sum[0] + cm[0]);
|
||||
col_sum[0] = tl::AllReduce<tl::SumOp, 16, 4, 0, tl::NamedBarrier<32>>::run(col_sum[0]);
|
||||
cm[0] = (cm[0] / (col_sum[0] + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
}
|
||||
if ((((int)threadIdx.x) >> 4) == 0) {
|
||||
comb_mix[((((int64_t)((int)blockIdx.x)) * (int64_t)16) + (((int64_t)((int)threadIdx.x)) & (int64_t)15))] = cm[0];
|
||||
}
|
||||
} else {
|
||||
if (((int)threadIdx.x) < 36) {
|
||||
((float*)buf_dyn_shmem)[(((int)threadIdx.x) + 7224)] = ((0x1p+0f/*1.000000e+00*/ / (0x1p+0f/*1.000000e+00*/ + expf((0x0p+0f/*0.000000e+00*/ - ((((float*)buf_dyn_shmem)[(((int)threadIdx.x) - 32)] * hc_scale[0]) + hc_base[(((int)threadIdx.x) - 32)]))))) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/);
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i = 0; i < 4; ++i) {
|
||||
float broadcast_var = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(sumsq_per_pos + (i * 4)) = make_float4(broadcast_var, broadcast_var, broadcast_var, broadcast_var);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_1 = 0; i_1 < 8; ++i_1) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_1 >> 1) * 1024) + (((((i_1 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_1) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_1) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8))])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_2 = 0; i_2 < 8; ++i_2) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i_2 >> 1) * 1024) + (((((i_2 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 4144)])), (&(residual[((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_2) >> (int64_t)1) * (int64_t)4096)) + (((((((int64_t)i_2) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)1024)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h = 0; i0_h < 2; ++i0_h) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_3 = 0; i_3 < 8; ++i_3) {
|
||||
*(uint4*)(xs_local_cast + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h * 4096) + (i_3 * 512)) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec = 0; vec < 2; ++vec) {
|
||||
float4 __2;
|
||||
uint2 v_ = *(uint2*)(xs_local_cast + (vec * 4));
|
||||
((float2*)(&__2))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[0]);
|
||||
((float2*)(&__2))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v_))[1]);
|
||||
*(float4*)(xl + ((i_3 * 8) + (vec * 4))) = __2;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_4 = 0; i_4 < 8; ++i_4) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h * 4096) + ((i_4 >> 1) * 1024)) + (((((i_4 * 64) + ((int)threadIdx.x)) + 96) & 127) * 8)) + 48)])), (&(residual[(((((((int64_t)((int)blockIdx.x)) * (int64_t)16384) + ((((int64_t)i_4) >> (int64_t)1) * (int64_t)4096)) + (((int64_t)i0_h) * (int64_t)1024)) + (((((((int64_t)i_4) * (int64_t)64) + ((int64_t)((int)threadIdx.x))) + (int64_t)96) & (int64_t)127) * (int64_t)8)) + (int64_t)2048)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_5 = 0; i_5 < 4; ++i_5) {
|
||||
float broadcast_var_1 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_5 * 4)) = make_float4(broadcast_var_1, broadcast_var_1, broadcast_var_1, broadcast_var_1);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
for (int i_hc = 0; i_hc < 4; ++i_hc) {
|
||||
float pre = ((float*)buf_dyn_shmem)[(i_hc + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_6 = 0; i_6 < 16; ++i_6) {
|
||||
ol[i_6] = (ol[i_6] + (pre * xl[((i_hc * 16) + i_6)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_7 = 0; i_7 < 4; ++i_7) {
|
||||
float4 __3;
|
||||
float4 v__1 = *(float4*)(sumsq_per_pos + (i_7 * 4));
|
||||
float4 __4;
|
||||
float4 v__2 = *(float4*)(ol + (i_7 * 4));
|
||||
__4.x = (v__2.x*v__2.x);
|
||||
__4.y = (v__2.y*v__2.y);
|
||||
__4.z = (v__2.z*v__2.z);
|
||||
__4.w = (v__2.w*v__2.w);
|
||||
__3.x = (v__1.x+__4.x);
|
||||
__3.y = (v__1.y+__4.y);
|
||||
__3.z = (v__1.z+__4.z);
|
||||
__3.w = (v__1.w+__4.w);
|
||||
*(float4*)(sumsq_per_pos + (i_7 * 4)) = __3;
|
||||
uint2 __5;
|
||||
float4 v__3 = *(float4*)(ol + (i_7 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[0] = __float22bfloat162_rn(((float2*)(&v__3))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__5))[1] = __float22bfloat162_rn(((float2*)(&v__3))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i0_h * 1024) + ((i_7 >> 1) * 512)) + (((int)threadIdx.x) * 8)) + ((i_7 & 1) * 4)) + 7984)) = __5;
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_8 = 0; i_8 < 8; ++i_8) {
|
||||
*(uint4*)(xs_local_cast_1 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_8 * 512) + (((int)threadIdx.x) * 8)) - 208));
|
||||
for (int vec_1 = 0; vec_1 < 2; ++vec_1) {
|
||||
float4 __6;
|
||||
uint2 v__4 = *(uint2*)(xs_local_cast_1 + (vec_1 * 4));
|
||||
((float2*)(&__6))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[0]);
|
||||
((float2*)(&__6))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__4))[1]);
|
||||
*(float4*)(xl + ((i_8 * 8) + (vec_1 * 4))) = __6;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_9 = 0; i_9 < 4; ++i_9) {
|
||||
float broadcast_var_2 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_9 * 4)) = make_float4(broadcast_var_2, broadcast_var_2, broadcast_var_2, broadcast_var_2);
|
||||
}
|
||||
for (int i_hc_1 = 0; i_hc_1 < 4; ++i_hc_1) {
|
||||
float pre_1 = ((float*)buf_dyn_shmem)[(i_hc_1 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_10 = 0; i_10 < 16; ++i_10) {
|
||||
ol[i_10] = (ol[i_10] + (pre_1 * xl[((i_hc_1 * 16) + i_10)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_11 = 0; i_11 < 4; ++i_11) {
|
||||
float4 __7;
|
||||
float4 v__5 = *(float4*)(sumsq_per_pos + (i_11 * 4));
|
||||
float4 __8;
|
||||
float4 v__6 = *(float4*)(ol + (i_11 * 4));
|
||||
__8.x = (v__6.x*v__6.x);
|
||||
__8.y = (v__6.y*v__6.y);
|
||||
__8.z = (v__6.z*v__6.z);
|
||||
__8.w = (v__6.w*v__6.w);
|
||||
__7.x = (v__5.x+__8.x);
|
||||
__7.y = (v__5.y+__8.y);
|
||||
__7.z = (v__5.z+__8.z);
|
||||
__7.w = (v__5.w+__8.w);
|
||||
*(float4*)(sumsq_per_pos + (i_11 * 4)) = __7;
|
||||
uint2 __9;
|
||||
float4 v__7 = *(float4*)(ol + (i_11 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[0] = __float22bfloat162_rn(((float2*)(&v__7))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__9))[1] = __float22bfloat162_rn(((float2*)(&v__7))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_11 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_11 & 1) * 4)) + 10032)) = __9;
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_12 = 0; i_12 < 8; ++i_12) {
|
||||
*(uint4*)(xs_local_cast_2 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_12 * 512) + (((int)threadIdx.x) * 8)) + 3888));
|
||||
for (int vec_2 = 0; vec_2 < 2; ++vec_2) {
|
||||
float4 __10;
|
||||
uint2 v__8 = *(uint2*)(xs_local_cast_2 + (vec_2 * 4));
|
||||
((float2*)(&__10))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[0]);
|
||||
((float2*)(&__10))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__8))[1]);
|
||||
*(float4*)(xl + ((i_12 * 8) + (vec_2 * 4))) = __10;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_13 = 0; i_13 < 4; ++i_13) {
|
||||
float broadcast_var_3 = 0x0p+0f/*0.000000e+00*/;
|
||||
*(float4*)(ol + (i_13 * 4)) = make_float4(broadcast_var_3, broadcast_var_3, broadcast_var_3, broadcast_var_3);
|
||||
}
|
||||
for (int i_hc_2 = 0; i_hc_2 < 4; ++i_hc_2) {
|
||||
float pre_2 = ((float*)buf_dyn_shmem)[(i_hc_2 + 7256)];
|
||||
#pragma unroll
|
||||
for (int i_14 = 0; i_14 < 16; ++i_14) {
|
||||
ol[i_14] = (ol[i_14] + (pre_2 * xl[((i_hc_2 * 16) + i_14)]));
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_15 = 0; i_15 < 4; ++i_15) {
|
||||
float4 __11;
|
||||
float4 v__9 = *(float4*)(sumsq_per_pos + (i_15 * 4));
|
||||
float4 __12;
|
||||
float4 v__10 = *(float4*)(ol + (i_15 * 4));
|
||||
__12.x = (v__10.x*v__10.x);
|
||||
__12.y = (v__10.y*v__10.y);
|
||||
__12.z = (v__10.z*v__10.z);
|
||||
__12.w = (v__10.w*v__10.w);
|
||||
__11.x = (v__9.x+__12.x);
|
||||
__11.y = (v__9.y+__12.y);
|
||||
__11.z = (v__9.z+__12.z);
|
||||
__11.w = (v__9.w+__12.w);
|
||||
*(float4*)(sumsq_per_pos + (i_15 * 4)) = __11;
|
||||
uint2 __13;
|
||||
float4 v__11 = *(float4*)(ol + (i_15 * 4));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[0] = __float22bfloat162_rn(((float2*)(&v__11))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__13))[1] = __float22bfloat162_rn(((float2*)(&v__11))[1]);
|
||||
*(uint2*)(((bfloat16_t*)buf_dyn_shmem) + (((((i_15 >> 1) * 512) + (((int)threadIdx.x) * 8)) + ((i_15 & 1) * 4)) + 11056)) = __13;
|
||||
}
|
||||
sumsq[0] = 0x0p+0f/*0.000000e+00*/;
|
||||
#pragma unroll
|
||||
for (int rv = 0; rv < 16; ++rv) {
|
||||
sumsq[0] = (sumsq[0] + sumsq_per_pos[(((rv & 1) * 8) + (rv >> 1))]);
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
sumsq[0] = tl::AllReduce<tl::SumOp, 64, 1, 32, tl::NamedBarrier<64>>::run(sumsq[0], (&(((float*)buf_dyn_shmem)[7192])));
|
||||
rsqrt_norm[0] = rsqrtf(((sumsq[0] / 0x1p+12f/*4.096000e+03*/) + 0x1.0c6f7a0b5ed8dp-20f/*1.000000e-06*/));
|
||||
#pragma unroll
|
||||
for (int i_16 = 0; i_16 < 2; ++i_16) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_16 * 512) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[(((i_16 * 512) + (((int)threadIdx.x) * 8)) - 256)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
#pragma unroll
|
||||
for (int i_17 = 0; i_17 < 2; ++i_17) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 13104)])), (&(norm_weight[(((i_17 * 512) + (((int)threadIdx.x) * 8)) + 768)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
for (int i0_h_1 = 0; i0_h_1 < 2; ++i0_h_1) {
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_18 = 0; i_18 < 2; ++i_18) {
|
||||
*(uint4*)(w_shared_local_cast_3 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_18 * 512)) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_3 = 0; vec_3 < 2; ++vec_3) {
|
||||
float4 __14;
|
||||
uint2 v__12 = *(uint2*)(w_shared_local_cast_3 + (vec_3 * 4));
|
||||
((float2*)(&__14))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[0]);
|
||||
((float2*)(&__14))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__12))[1]);
|
||||
*(float4*)(w_local + ((i_18 * 8) + (vec_3 * 4))) = __14;
|
||||
}
|
||||
}
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_19 = 0; i_19 < 2; ++i_19) {
|
||||
tl::cp_async_gs<16>((&(((bfloat16_t*)buf_dyn_shmem)[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 12080)])), (&(norm_weight[((((i0_h_1 * 1024) + (i_19 * 512)) + (((int)threadIdx.x) * 8)) + 1792)])));
|
||||
}
|
||||
tl::cp_async_commit();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_20 = 0; i_20 < 2; ++i_20) {
|
||||
*(uint4*)(output_shared_local_cast_4 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + ((((i0_h_1 * 1024) + (i_20 * 512)) + (((int)threadIdx.x) * 8)) + 7984));
|
||||
for (int vec_4 = 0; vec_4 < 2; ++vec_4) {
|
||||
float4 __15;
|
||||
float4 __16;
|
||||
float4 __17;
|
||||
uint2 v__13 = *(uint2*)(output_shared_local_cast_4 + (vec_4 * 4));
|
||||
((float2*)(&__17))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[0]);
|
||||
((float2*)(&__17))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__13))[1]);
|
||||
float4 v__14 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__16.x = (__17.x*v__14.x);
|
||||
__16.y = (__17.y*v__14.y);
|
||||
__16.z = (__17.z*v__14.z);
|
||||
__16.w = (__17.w*v__14.w);
|
||||
float4 v__15 = *(float4*)(w_local + ((i_20 * 8) + (vec_4 * 4)));
|
||||
__15.x = (__16.x*v__15.x);
|
||||
__15.y = (__16.y*v__15.y);
|
||||
__15.z = (__16.z*v__15.z);
|
||||
__15.w = (__16.w*v__15.w);
|
||||
*(float4*)(ol_1 + ((i_20 * 8) + (vec_4 * 4))) = __15;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_21 = 0; i_21 < 2; ++i_21) {
|
||||
for (int vec_5 = 0; vec_5 < 2; ++vec_5) {
|
||||
uint2 __18;
|
||||
float4 v__16 = *(float4*)(ol_1 + ((i_21 * 8) + (vec_5 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[0] = __float22bfloat162_rn(((float2*)(&v__16))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__18))[1] = __float22bfloat162_rn(((float2*)(&v__16))[1]);
|
||||
*(uint2*)(layer_input_local_cast_5 + (vec_5 * 4)) = __18;
|
||||
}
|
||||
*(uint4*)(layer_input + (((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i0_h_1) * (int64_t)1024)) + (((int64_t)i_21) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) - (int64_t)256)) = *(uint4*)(layer_input_local_cast_5 + 0);
|
||||
}
|
||||
}
|
||||
tl::cp_async_wait<1>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_22 = 0; i_22 < 2; ++i_22) {
|
||||
*(uint4*)(w_shared_local_cast_6 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_22 * 512) + (((int)threadIdx.x) * 8)) + 12080));
|
||||
for (int vec_6 = 0; vec_6 < 2; ++vec_6) {
|
||||
float4 __19;
|
||||
uint2 v__17 = *(uint2*)(w_shared_local_cast_6 + (vec_6 * 4));
|
||||
((float2*)(&__19))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[0]);
|
||||
((float2*)(&__19))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__17))[1]);
|
||||
*(float4*)(w_local + ((i_22 * 8) + (vec_6 * 4))) = __19;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_23 = 0; i_23 < 2; ++i_23) {
|
||||
*(uint4*)(output_shared_local_cast_7 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_23 * 512) + (((int)threadIdx.x) * 8)) + 10032));
|
||||
for (int vec_7 = 0; vec_7 < 2; ++vec_7) {
|
||||
float4 __20;
|
||||
float4 __21;
|
||||
float4 __22;
|
||||
uint2 v__18 = *(uint2*)(output_shared_local_cast_7 + (vec_7 * 4));
|
||||
((float2*)(&__22))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[0]);
|
||||
((float2*)(&__22))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__18))[1]);
|
||||
float4 v__19 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__21.x = (__22.x*v__19.x);
|
||||
__21.y = (__22.y*v__19.y);
|
||||
__21.z = (__22.z*v__19.z);
|
||||
__21.w = (__22.w*v__19.w);
|
||||
float4 v__20 = *(float4*)(w_local + ((i_23 * 8) + (vec_7 * 4)));
|
||||
__20.x = (__21.x*v__20.x);
|
||||
__20.y = (__21.y*v__20.y);
|
||||
__20.z = (__21.z*v__20.z);
|
||||
__20.w = (__21.w*v__20.w);
|
||||
*(float4*)(ol_1 + ((i_23 * 8) + (vec_7 * 4))) = __20;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_24 = 0; i_24 < 2; ++i_24) {
|
||||
for (int vec_8 = 0; vec_8 < 2; ++vec_8) {
|
||||
uint2 __23;
|
||||
float4 v__21 = *(float4*)(ol_1 + ((i_24 * 8) + (vec_8 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[0] = __float22bfloat162_rn(((float2*)(&v__21))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__23))[1] = __float22bfloat162_rn(((float2*)(&v__21))[1]);
|
||||
*(uint2*)(layer_input_local_cast_8 + (vec_8 * 4)) = __23;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_24) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)1792)) = *(uint4*)(layer_input_local_cast_8 + 0);
|
||||
}
|
||||
tl::cp_async_wait<0>();
|
||||
tl::__sync_thread_partial<3, 64>();
|
||||
#pragma unroll
|
||||
for (int i_25 = 0; i_25 < 2; ++i_25) {
|
||||
*(uint4*)(w_shared_local_cast_9 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_25 * 512) + (((int)threadIdx.x) * 8)) + 13104));
|
||||
for (int vec_9 = 0; vec_9 < 2; ++vec_9) {
|
||||
float4 __24;
|
||||
uint2 v__22 = *(uint2*)(w_shared_local_cast_9 + (vec_9 * 4));
|
||||
((float2*)(&__24))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[0]);
|
||||
((float2*)(&__24))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__22))[1]);
|
||||
*(float4*)(w_local + ((i_25 * 8) + (vec_9 * 4))) = __24;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_26 = 0; i_26 < 2; ++i_26) {
|
||||
*(uint4*)(output_shared_local_cast_10 + 0) = *(uint4*)(((bfloat16_t*)buf_dyn_shmem) + (((i_26 * 512) + (((int)threadIdx.x) * 8)) + 11056));
|
||||
for (int vec_10 = 0; vec_10 < 2; ++vec_10) {
|
||||
float4 __25;
|
||||
float4 __26;
|
||||
float4 __27;
|
||||
uint2 v__23 = *(uint2*)(output_shared_local_cast_10 + (vec_10 * 4));
|
||||
((float2*)(&__27))[0] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[0]);
|
||||
((float2*)(&__27))[1] = __bfloat1622float2((reinterpret_cast<__nv_bfloat162*>(&v__23))[1]);
|
||||
float4 v__24 = make_float4(rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0], rsqrt_norm[0]);
|
||||
__26.x = (__27.x*v__24.x);
|
||||
__26.y = (__27.y*v__24.y);
|
||||
__26.z = (__27.z*v__24.z);
|
||||
__26.w = (__27.w*v__24.w);
|
||||
float4 v__25 = *(float4*)(w_local + ((i_26 * 8) + (vec_10 * 4)));
|
||||
__25.x = (__26.x*v__25.x);
|
||||
__25.y = (__26.y*v__25.y);
|
||||
__25.z = (__26.z*v__25.z);
|
||||
__25.w = (__26.w*v__25.w);
|
||||
*(float4*)(ol_1 + ((i_26 * 8) + (vec_10 * 4))) = __25;
|
||||
}
|
||||
}
|
||||
#pragma unroll
|
||||
for (int i_27 = 0; i_27 < 2; ++i_27) {
|
||||
for (int vec_11 = 0; vec_11 < 2; ++vec_11) {
|
||||
uint2 __28;
|
||||
float4 v__26 = *(float4*)(ol_1 + ((i_27 * 8) + (vec_11 * 4)));
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[0] = __float22bfloat162_rn(((float2*)(&v__26))[0]);
|
||||
(reinterpret_cast<__nv_bfloat162*>(&__28))[1] = __float22bfloat162_rn(((float2*)(&v__26))[1]);
|
||||
*(uint2*)(layer_input_local_cast_11 + (vec_11 * 4)) = __28;
|
||||
}
|
||||
*(uint4*)(layer_input + ((((((int64_t)((int)blockIdx.x)) * (int64_t)4096) + (((int64_t)i_27) * (int64_t)512)) + (((int64_t)((int)threadIdx.x)) * (int64_t)8)) + (int64_t)2816)) = *(uint4*)(layer_input_local_cast_11 + 0);
|
||||
}
|
||||
}
|
||||
cudaTriggerProgrammaticLaunchCompletion();
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
Loading…
x
Reference in New Issue
Block a user