MCPcopy Create free account
hub / github.com/cactus-compute/cactus / test_embedding_operation

Function test_embedding_operation

tests/test_graph.cpp:1216–1297  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1214}
1215
1216bool 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

Callers 1

mainFunction · 0.85

Calls 9

sizeMethod · 0.80
dataMethod · 0.80
inputMethod · 0.45
set_inputMethod · 0.45
set_grouped_scalesMethod · 0.45
set_interleavedMethod · 0.45
embeddingMethod · 0.45
executeMethod · 0.45
get_outputMethod · 0.45

Tested by

no test coverage detected