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

Function maxpooling_grad

imperative/python/megengine/xla/rules/nn.py:501–544  ·  view source on GitHub ↗
(
    x,
    dy,
    kernel,
    stride,
    padding,
    base_dilation=None,
    kernel_dilation=None,
    expand_padding=True,
)

Source from the content-addressed store, hash-verified

499
500
501def maxpooling_grad(
502 x,
503 dy,
504 kernel,
505 stride,
506 padding,
507 base_dilation=None,
508 kernel_dilation=None,
509 expand_padding=True,
510):
511 assert base_dilation is None and kernel_dilation is None
512 assert expand_padding == True
513 padding = [(p, p) if isinstance(p, int) else p for p in padding]
514 dxdtype, dxshape = x.dtype, x.shape
515 assert dxdtype == "float32" or dxdtype == "float16"
516
517 org_padding, new_padding = padding, padding
518 if expand_padding:
519 pads = [(lo, hi, 0) for (lo, hi) in padding]
520 padded_x = pad(x, _get_max_identity(dxdtype), pads)
521 new_padding = [(0, 0) for _ in padding]
522
523 selector = lambda x, y: x >= y
524 scatter = lambda x, y: x + y
525 out = _select_and_scatter(
526 padded_x,
527 dy,
528 np.array(0.0, dtype=dy.dtype),
529 kernel,
530 stride,
531 new_padding,
532 selector,
533 scatter,
534 )
535
536 if expand_padding:
537 start_indices = [lo for (lo, hi) in org_padding]
538 stop_indices = [lo + d for ((lo, hi), d) in zip(org_padding, dxshape)]
539 slices = [
540 slice(start, stop, 1) for start, stop in zip(start_indices, stop_indices)
541 ]
542 out = index_with_slices(out, slices)
543
544 return out
545
546
547def avgpooling_grad(

Callers 2

pooling_backward_lowerFunction · 0.85

Calls 5

_get_max_identityFunction · 0.85
_select_and_scatterFunction · 0.85
index_with_slicesFunction · 0.85
arrayMethod · 0.80
padFunction · 0.70

Tested by

no test coverage detected