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

Method embedding

cactus/graph/graph_builder.cpp:1440–1458  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1438}
1439
1440size_t CactusGraph::embedding(size_t embedding_tensor, size_t indices) {
1441 const auto& emb_buffer = get_output_buffer(embedding_tensor);
1442 const auto& idx_shape = get_output_buffer(indices).shape;
1443
1444 if (emb_buffer.shape.size() != 2) {
1445 std::cerr << "Error: Embedding tensor " << embedding_tensor << " has invalid shape: [";
1446 for(auto d : emb_buffer.shape) std::cerr << d << ",";
1447 std::cerr << "]. OpType=" << (int)nodes_[node_index_map_[embedding_tensor]]->op_type
1448 << " ExtData=" << emb_buffer.external_data << std::endl;
1449 throw std::runtime_error("Embedding tensor must be 2D [vocab_size, hidden_dim]");
1450 }
1451
1452 std::vector<size_t> output_shape = idx_shape;
1453 output_shape.push_back(emb_buffer.shape[1]);
1454
1455 OpParams params;
1456 params.output_precision = Precision::FP16;
1457 return add_node(OpType::EMBEDDING, {embedding_tensor, indices}, output_shape, params);
1458}
1459
1460size_t CactusGraph::bilinear_interpolation(size_t pos_embeds, size_t dst_height, size_t dst_width, bool align_corners) {
1461 const auto& pos_embeds_buffer = get_output_buffer(pos_embeds);

Callers

nothing calls this directly

Calls 1

sizeMethod · 0.80

Tested by

no test coverage detected