| 1438 | } |
| 1439 | |
| 1440 | size_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 | |
| 1460 | size_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); |