pulsatrix
Loading...
Searching...
No Matches
group_norm_module.hpp File Reference

Group normalization (Wu & He, 2018) – rank-3 (C, H, W), matching Conv2DModule's convention, unlike RMSNormModule/LayerNormModule's rank-1 feature-vector scope. More...

#include <optional>
#include <initializer_list>
#include <vector>
#include "pulsatrix/module.hpp"
Include dependency graph for group_norm_module.hpp:

Go to the source code of this file.

Classes

class  pulsatrix::GroupNormModule
 Splits num_channels into num_groups equal-size groups; each group's mean/std is computed over every (channel-in-group, H, W) element jointly, per batch row n, then y_{n,c,h,w} = gamma_c * (x_{n,c,h,w} - mu_{n,g})/std_{n,g} + beta_c, gamma/beta per-channel (shape (num_channels,), not per-group, not per-batch-row). Batched – input/output are rank-4 (N, channels, H, W), migrated from the original unbatched (rank-3) scope by campaign_exai_dl_library_batch_dimension_support, matching Conv2DModule's own (not-yet-migrated) rank-3 convention plus a leading batch dim. H/W are not fixed at construction (only num_groups/num_channels are), so this module accepts any spatial size at forward() time, exactly like Conv2DModule does. More...
 

Namespaces

namespace  pulsatrix
 

Detailed Description

Group normalization (Wu & He, 2018) – rank-3 (C, H, W), matching Conv2DModule's convention, unlike RMSNormModule/LayerNormModule's rank-1 feature-vector scope.