Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- import jax
- import jax.numpy as jnp
- import flax.linen as nn
- PARAM_DTYPE = jnp.float32
- COMPUTE_DTYPE = jnp.bfloat16
- def normalized_smish(x):
- """Normalized version of x * tanh(log(1 + sigmoid(x)))."""
- smish_mean = 0.120057022660301
- smish_stddev = 0.379395531405842
- sigmoid_x = jax.nn.sigmoid(x)
- x = x * jnp.tanh(jnp.log1p(sigmoid_x))
- x = (x - smish_mean) / smish_stddev
- return x
- class IdentityGate(nn.Module):
- """Control: do nothing."""
- def __call__(self, x):
- return x
- class SqueezeExcitationGate(nn.Module):
- """Standard squeeze-and-excitation with reduction ratio 8."""
- reduction_ratio: int = 8
- @nn.compact
- def __call__(self, x):
- batch, _height, _width, channels = x.shape
- hidden_features = channels // self.reduction_ratio
- channel_mean = jnp.mean(
- x,
- axis=(1, 2),
- dtype=jnp.float32,
- )
- hidden = nn.Dense(
- features=hidden_features,
- kernel_init=nn.initializers.kaiming_normal(),
- bias_init=nn.initializers.zeros,
- param_dtype=PARAM_DTYPE,
- dtype=COMPUTE_DTYPE,
- )(channel_mean)
- hidden = normalized_smish(hidden)
- logits = nn.Dense(
- features=channels,
- kernel_init=nn.initializers.zeros,
- bias_init=nn.initializers.zeros,
- param_dtype=PARAM_DTYPE,
- dtype=COMPUTE_DTYPE,
- )(hidden)
- scale = (
- 2.0 * jax.nn.sigmoid(logits.astype(jnp.float32))
- ).astype(COMPUTE_DTYPE)
- x = x * scale.reshape(batch, 1, 1, channels)
- return x
- class EfficientChannelAttentionGate(nn.Module):
- """ECA with a shared 1D channel kernel and no bias."""
- kernel_size: int
- @nn.compact
- def __call__(self, x):
- batch, _height, _width, channels = x.shape
- channel_mean = jnp.mean(
- x,
- axis=(1, 2),
- dtype=jnp.float32,
- )
- logits = nn.Conv(
- features=1,
- kernel_size=(self.kernel_size,),
- padding="SAME", # Zero padding, as used in these experiments.
- use_bias=False,
- kernel_init=nn.initializers.zeros,
- param_dtype=PARAM_DTYPE,
- dtype=COMPUTE_DTYPE,
- )(channel_mean[..., None]).squeeze(-1)
- scale = (
- 2.0 * jax.nn.sigmoid(logits.astype(jnp.float32))
- ).astype(COMPUTE_DTYPE)
- x = x * scale.reshape(batch, 1, 1, channels)
- return x
- # The two ECA rows used the same gate with different kernel sizes:
- # EfficientChannelAttentionGate(kernel_size=3)
- # EfficientChannelAttentionGate(kernel_size=1)
- class CenterMaskedEfficientChannelAttentionGate(nn.Module):
- """ECA3 with its centre kernel weight permanently masked out."""
- kernel_size: int = 3
- @nn.compact
- def __call__(self, x):
- batch, _height, _width, channels = x.shape
- channel_mean = jnp.mean(
- x,
- axis=(1, 2),
- dtype=jnp.float32,
- )
- kernel = self.param(
- "kernel",
- nn.initializers.zeros,
- (self.kernel_size, 1, 1),
- PARAM_DTYPE,
- )
- # For kernel size 3, this produces the mask [1, 0, 1].
- kernel_mask = jnp.ones_like(kernel)
- kernel_mask = kernel_mask.at[
- self.kernel_size // 2, :, :
- ].set(0)
- kernel = kernel * kernel_mask
- logits = jax.lax.conv_general_dilated(
- channel_mean.astype(COMPUTE_DTYPE)[..., None],
- kernel.astype(COMPUTE_DTYPE),
- window_strides=(1,),
- padding="SAME", # Zero padding along the channel axis.
- dimension_numbers=("NWC", "WIO", "NWC"),
- ).squeeze(-1)
- scale = (
- 2.0 * jax.nn.sigmoid(logits.astype(jnp.float32))
- ).astype(COMPUTE_DTYPE)
- x = x * scale.reshape(batch, 1, 1, channels)
- return x
- class PerChannelGate(nn.Module):
- """One independent learned weight for each channel."""
- @nn.compact
- def __call__(self, x):
- batch, _height, _width, channels = x.shape
- channel_mean = jnp.mean(
- x,
- axis=(1, 2),
- dtype=jnp.float32,
- )
- channel_weights = self.param(
- "channel_weights",
- nn.initializers.zeros,
- (channels,),
- PARAM_DTYPE,
- )
- logits = (
- channel_mean *
- channel_weights.astype(jnp.float32)
- )
- scale = (
- 2.0 * jax.nn.sigmoid(logits)
- ).astype(COMPUTE_DTYPE)
- x = x * scale.reshape(batch, 1, 1, channels)
- return x
Advertisement
Add Comment
Please, Sign In to add comment