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

Function compute_stats_pool_node

cactus/graph/graph_ops_nn.cpp:2276–2305  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

2274}
2275
2276void compute_stats_pool_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes,
2277 const std::unordered_map<size_t, size_t>& node_index_map) {
2278 const auto& input = get_input(node, 0, nodes, node_index_map);
2279 const __fp16* src = input.data_as<__fp16>();
2280 __fp16* dst = node.output_buffer.data_as<__fp16>();
2281
2282 size_t batch = input.shape[0];
2283 size_t total_per_batch = input.total_size / batch;
2284 size_t T = input.shape.back();
2285 size_t features = total_per_batch / T;
2286
2287 for (size_t b = 0; b < batch; ++b) {
2288 const __fp16* batch_src = src + b * total_per_batch;
2289 __fp16* batch_dst = dst + b * features * 2;
2290
2291 for (size_t f = 0; f < features; ++f) {
2292 float sum = 0.0f, sum_sq = 0.0f;
2293 for (size_t t = 0; t < T; ++t) {
2294 float v = static_cast<float>(batch_src[f * T + t]);
2295 sum += v;
2296 sum_sq += v * v;
2297 }
2298 float mean = sum / static_cast<float>(T);
2299 float var = T > 1 ? (sum_sq - static_cast<float>(T) * mean * mean) / static_cast<float>(T - 1) : 0.0f;
2300 float std_val = sqrtf(fmaxf(var, 0.0f));
2301 batch_dst[f] = static_cast<__fp16>(mean);
2302 batch_dst[features + f] = static_cast<__fp16>(std_val);
2303 }
2304 }
2305}
2306
2307void compute_weighted_stats_pool_node(GraphNode& node, const std::vector<std::unique_ptr<GraphNode>>& nodes,
2308 const std::unordered_map<size_t, size_t>& node_index_map) {

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected