Normalization Factory
Normalization factory.
Selects the appropriate normalization module from the model configuration,
mirroring the optimizer factory pattern in src/model/optimizer/factory.py.
- src.model.norm.factory.get_norm(config)[source]
Factory function that returns a normalization module based on config.
Selects among
LayerNorm,DynamicTanhNorm,Derf,RMSNorm,pRMSNorm(partial RMSNorm), andFlashNorm(weightless RMSNorm with optional partial-RMS composition) based on thenorm_typefield in the configuration.- Parameters:
config – Model configuration object with attributes
norm_type(one of"layer_norm","dynamic_tanh","derf","rms_norm","prms_norm","flash_norm") andhidden_size. Whennorm_type == "prms_norm", the optionalprms_partial_ratioattribute (default0.0625) controls the fraction of dimensions used for RMS estimation. Whennorm_type == "flash_norm", the optionalflashnorm_partial_ratioattribute (default0.0) activates the partial-RMS variant of FlashNorm.- Returns:
A normalization
nn.Moduleinstance appropriate for the requestednorm_type.- Raises:
AttributeError – If
configdoes not havenorm_typeorhidden_sizeattributes.