Download local_response_norm.py from PrakhAI/AIPlane2: direct link, hf CLI and curl.
- Browser
- Download file 313 Bytes
-
https://huggingface.co/spaces/PrakhAI/AIPlane2/resolve/main/local_response_norm.py
- Command line
-
hf download hf://spaces/PrakhAI/AIPlane2/local_response_norm.py
-
curl -L -o local_response_norm.py https://huggingface.co/spaces/PrakhAI/AIPlane2/resolve/main/local_response_norm.py
313 Bytes
| from flax import linen as nn | |
| import jax | |
| import jax.numpy as jnp | |
| class LocalResponseNorm(nn.Module): | |
| def __call__( | |
| self, | |
| value: jax.Array | |
| ) -> jax.Array: | |
| return value / jnp.repeat(jnp.expand_dims((1e-8 + (value**2).mean(axis=-1))**0.5, axis=-1), repeats=value.shape[-1], axis=-1) |