| 1214 | } |
| 1215 | |
| 1216 | bool test_embedding_operation() { |
| 1217 | CactusGraph graph; |
| 1218 | |
| 1219 | const size_t vocab_size = 4; |
| 1220 | const size_t hidden_dim = 8; |
| 1221 | const size_t group_size = 8; |
| 1222 | const size_t num_groups = hidden_dim / group_size; |
| 1223 | const size_t BLOCK_SIZE = 4; |
| 1224 | |
| 1225 | std::vector<int8_t> emb_rowmajor(vocab_size * hidden_dim); |
| 1226 | for (size_t row = 0; row < vocab_size; ++row) { |
| 1227 | for (size_t k = 0; k < hidden_dim; ++k) { |
| 1228 | emb_rowmajor[row * hidden_dim + k] = static_cast<int8_t>((row + 1) * 10 + k); |
| 1229 | } |
| 1230 | } |
| 1231 | |
| 1232 | std::vector<int8_t> emb_interleaved(vocab_size * hidden_dim); |
| 1233 | size_t N_blocks = vocab_size / BLOCK_SIZE; |
| 1234 | size_t K_groups = hidden_dim / BLOCK_SIZE; |
| 1235 | |
| 1236 | for (size_t n_blk = 0; n_blk < N_blocks; ++n_blk) { |
| 1237 | for (size_t k_grp = 0; k_grp < K_groups; ++k_grp) { |
| 1238 | for (size_t lane = 0; lane < BLOCK_SIZE; ++lane) { |
| 1239 | for (size_t k_within = 0; k_within < BLOCK_SIZE; ++k_within) { |
| 1240 | size_t src_row = n_blk * BLOCK_SIZE + lane; |
| 1241 | size_t src_k = k_grp * BLOCK_SIZE + k_within; |
| 1242 | size_t dst_idx = (n_blk * K_groups + k_grp) * 16 + lane * 4 + k_within; |
| 1243 | emb_interleaved[dst_idx] = emb_rowmajor[src_row * hidden_dim + src_k]; |
| 1244 | } |
| 1245 | } |
| 1246 | } |
| 1247 | } |
| 1248 | |
| 1249 | std::vector<__fp16> scales_rowmajor(vocab_size * num_groups); |
| 1250 | for (size_t i = 0; i < scales_rowmajor.size(); ++i) { |
| 1251 | scales_rowmajor[i] = static_cast<__fp16>(1.0f); |
| 1252 | } |
| 1253 | |
| 1254 | std::vector<__fp16> scales_interleaved(vocab_size * num_groups); |
| 1255 | for (size_t n_blk = 0; n_blk < N_blocks; ++n_blk) { |
| 1256 | for (size_t g = 0; g < num_groups; ++g) { |
| 1257 | for (size_t lane = 0; lane < BLOCK_SIZE; ++lane) { |
| 1258 | size_t src_row = n_blk * BLOCK_SIZE + lane; |
| 1259 | size_t dst_idx = (n_blk * num_groups + g) * BLOCK_SIZE + lane; |
| 1260 | scales_interleaved[dst_idx] = scales_rowmajor[src_row * num_groups + g]; |
| 1261 | } |
| 1262 | } |
| 1263 | } |
| 1264 | |
| 1265 | size_t embeddings = graph.input({vocab_size, hidden_dim}, Precision::INT8); |
| 1266 | graph.set_input(embeddings, emb_interleaved.data(), Precision::INT8); |
| 1267 | |
| 1268 | graph.set_grouped_scales(embeddings, group_size, num_groups, scales_interleaved.data()); |
| 1269 | graph.set_interleaved(embeddings, true, vocab_size); |
| 1270 | |
| 1271 | size_t indices = graph.input({4}, Precision::INT8); |
| 1272 | size_t embedded = graph.embedding(embeddings, indices); |
| 1273 |
no test coverage detected