| 1407 | |
| 1408 | @subgraph("SyncBnStage2", dtype, device, 7) |
| 1409 | def syncbn_stage2(inputs, f, c): |
| 1410 | running_mean, running_var, momentum = inputs[0:3] |
| 1411 | reduce_size, channel_x1s, channel_x2s, channel_mean = inputs[3:7] |
| 1412 | c1_minus_momentum = f("-", c(1), momentum) |
| 1413 | reduce_size_minus_c1 = f("-", reduce_size, c(1)) |
| 1414 | running_mean = f( |
| 1415 | "fma4", running_mean, momentum, c1_minus_momentum, channel_mean, |
| 1416 | ) |
| 1417 | channel_variance_unbiased = f( |
| 1418 | "+", |
| 1419 | f( |
| 1420 | "/", |
| 1421 | f("**", channel_x1s, c(2)), |
| 1422 | f("*", f("-", reduce_size), reduce_size_minus_c1), |
| 1423 | ), |
| 1424 | f("/", channel_x2s, reduce_size_minus_c1), |
| 1425 | ) |
| 1426 | running_var = f( |
| 1427 | "fma4", running_var, momentum, c1_minus_momentum, channel_variance_unbiased |
| 1428 | ) |
| 1429 | return (running_mean, running_var), (True, True) |
| 1430 | |
| 1431 | @subgraph("SyncBnConcatStats", dtype, device, 3) |
| 1432 | def syncbn_concat_stats(inputs, f, c): |