Download code/kernels/sp_input/u8_to_bf16.cpp from changh95/superpoint-p150: direct link, hf CLI and curl.
- Browser
- Download file 2.43 kB
-
https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_input/u8_to_bf16.cpp
- Command line
-
hf download hf://changh95/superpoint-p150/code/kernels/sp_input/u8_to_bf16.cpp
-
curl -L -o u8_to_bf16.cpp https://huggingface.co/changh95/superpoint-p150/resolve/main/code/kernels/sp_input/u8_to_bf16.cpp
2.43 kB
| // SPDX-FileCopyrightText: © 2026 Tenstorrent USA, Inc. | |
| // SPDX-License-Identifier: Apache-2.0 | |
| // uint8 image shard -> bf16 shard through a 256-entry table (bf16(fp32(u) / 255), built on the host: | |
| // bit-identical to the host fp32 /255 + bf16 cast). Both data-movement RISCs, each one half of the | |
| // core's pixels, local L1 only (input and output are height-sharded on the same cores, same pixel split). | |
| // CT args: cb_in, cb_out, n_px (pixels per core, multiple of 8), proc (0 BRISC / 1 NCRISC) | |
| // The table is compiled in: the host prepends `#define SP_U8_LUT <256 comma-separated bf16 bit patterns>` | |
| // (no runtime args, nothing to dispatch per replay). | |
| static const uint16_t lut[256] = {SP_U8_LUT}; | |
| void kernel_main() { | |
| constexpr uint32_t cb_in = get_compile_time_arg_val(0); | |
| constexpr uint32_t cb_out = get_compile_time_arg_val(1); | |
| constexpr uint32_t n_px = get_compile_time_arg_val(2); | |
| constexpr uint32_t proc = get_compile_time_arg_val(3); | |
| constexpr uint32_t n_words = n_px / 4; // 4 pixels per input word | |
| constexpr uint32_t half = (n_words / 2) & ~3u; | |
| constexpr uint32_t w0 = proc == 0 ? 0 : half; | |
| constexpr uint32_t w1 = proc == 0 ? half : n_words; | |
| volatile tt_l1_ptr uint32_t* in = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_read_ptr(cb_in)); | |
| volatile tt_l1_ptr uint32_t* out = reinterpret_cast<volatile tt_l1_ptr uint32_t*>(get_write_ptr(cb_out)); | |
| static_assert((w1 - w0) % 4 == 0, "4-word unroll"); | |
| for (uint32_t w = w0; w < w1; w += 4) { | |
| // issue the four L1 loads back to back, then the table lookups / stores | |
| uint32_t a = in[w], b = in[w + 1], c = in[w + 2], d = in[w + 3]; | |
| volatile tt_l1_ptr uint32_t* o = out + 2 * w; | |
| o[0] = (uint32_t)lut[a & 0xFF] | ((uint32_t)lut[(a >> 8) & 0xFF] << 16); | |
| o[1] = (uint32_t)lut[(a >> 16) & 0xFF] | ((uint32_t)lut[a >> 24] << 16); | |
| o[2] = (uint32_t)lut[b & 0xFF] | ((uint32_t)lut[(b >> 8) & 0xFF] << 16); | |
| o[3] = (uint32_t)lut[(b >> 16) & 0xFF] | ((uint32_t)lut[b >> 24] << 16); | |
| o[4] = (uint32_t)lut[c & 0xFF] | ((uint32_t)lut[(c >> 8) & 0xFF] << 16); | |
| o[5] = (uint32_t)lut[(c >> 16) & 0xFF] | ((uint32_t)lut[c >> 24] << 16); | |
| o[6] = (uint32_t)lut[d & 0xFF] | ((uint32_t)lut[(d >> 8) & 0xFF] << 16); | |
| o[7] = (uint32_t)lut[(d >> 16) & 0xFF] | ((uint32_t)lut[d >> 24] << 16); | |
| } | |
| } | |