MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / syncbn_stage2

Function syncbn_stage2

imperative/python/megengine/functional/nn.py:1409–1429  ·  view source on GitHub ↗
(inputs, f, c)

Source from the content-addressed store, hash-verified

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):

Callers 1

sync_batch_normFunction · 0.85

Calls 1

fFunction · 0.50

Tested by

no test coverage detected