Guest User

Untitled

a guest
Aug 16th, 2026
143
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 4.78 KB | None | 0 0
  1. import jax
  2. import jax.numpy as jnp
  3. import flax.linen as nn
  4.  
  5.  
  6. PARAM_DTYPE = jnp.float32
  7. COMPUTE_DTYPE = jnp.bfloat16
  8.  
  9.  
  10. def normalized_smish(x):
  11.     """Normalized version of x * tanh(log(1 + sigmoid(x)))."""
  12.     smish_mean = 0.120057022660301
  13.     smish_stddev = 0.379395531405842
  14.  
  15.     sigmoid_x = jax.nn.sigmoid(x)
  16.     x = x * jnp.tanh(jnp.log1p(sigmoid_x))
  17.     x = (x - smish_mean) / smish_stddev
  18.     return x
  19.  
  20.  
  21. class IdentityGate(nn.Module):
  22.     """Control: do nothing."""
  23.  
  24.     def __call__(self, x):
  25.         return x
  26.  
  27.  
  28. class SqueezeExcitationGate(nn.Module):
  29.     """Standard squeeze-and-excitation with reduction ratio 8."""
  30.  
  31.     reduction_ratio: int = 8
  32.  
  33.     @nn.compact
  34.     def __call__(self, x):
  35.         batch, _height, _width, channels = x.shape
  36.         hidden_features = channels // self.reduction_ratio
  37.  
  38.         channel_mean = jnp.mean(
  39.             x,
  40.             axis=(1, 2),
  41.             dtype=jnp.float32,
  42.         )
  43.  
  44.         hidden = nn.Dense(
  45.             features=hidden_features,
  46.             kernel_init=nn.initializers.kaiming_normal(),
  47.             bias_init=nn.initializers.zeros,
  48.             param_dtype=PARAM_DTYPE,
  49.             dtype=COMPUTE_DTYPE,
  50.         )(channel_mean)
  51.         hidden = normalized_smish(hidden)
  52.  
  53.         logits = nn.Dense(
  54.             features=channels,
  55.             kernel_init=nn.initializers.zeros,
  56.             bias_init=nn.initializers.zeros,
  57.             param_dtype=PARAM_DTYPE,
  58.             dtype=COMPUTE_DTYPE,
  59.         )(hidden)
  60.  
  61.         scale = (
  62.             2.0 * jax.nn.sigmoid(logits.astype(jnp.float32))
  63.         ).astype(COMPUTE_DTYPE)
  64.  
  65.         x = x * scale.reshape(batch, 1, 1, channels)
  66.         return x
  67.  
  68.  
  69. class EfficientChannelAttentionGate(nn.Module):
  70.     """ECA with a shared 1D channel kernel and no bias."""
  71.  
  72.     kernel_size: int
  73.  
  74.     @nn.compact
  75.     def __call__(self, x):
  76.         batch, _height, _width, channels = x.shape
  77.  
  78.         channel_mean = jnp.mean(
  79.             x,
  80.             axis=(1, 2),
  81.             dtype=jnp.float32,
  82.         )
  83.  
  84.         logits = nn.Conv(
  85.             features=1,
  86.             kernel_size=(self.kernel_size,),
  87.             padding="SAME",  # Zero padding, as used in these experiments.
  88.             use_bias=False,
  89.             kernel_init=nn.initializers.zeros,
  90.             param_dtype=PARAM_DTYPE,
  91.             dtype=COMPUTE_DTYPE,
  92.         )(channel_mean[..., None]).squeeze(-1)
  93.  
  94.         scale = (
  95.             2.0 * jax.nn.sigmoid(logits.astype(jnp.float32))
  96.         ).astype(COMPUTE_DTYPE)
  97.  
  98.         x = x * scale.reshape(batch, 1, 1, channels)
  99.         return x
  100.  
  101.  
  102. # The two ECA rows used the same gate with different kernel sizes:
  103. # EfficientChannelAttentionGate(kernel_size=3)
  104. # EfficientChannelAttentionGate(kernel_size=1)
  105.  
  106.  
  107. class CenterMaskedEfficientChannelAttentionGate(nn.Module):
  108.     """ECA3 with its centre kernel weight permanently masked out."""
  109.  
  110.     kernel_size: int = 3
  111.  
  112.     @nn.compact
  113.     def __call__(self, x):
  114.         batch, _height, _width, channels = x.shape
  115.  
  116.         channel_mean = jnp.mean(
  117.             x,
  118.             axis=(1, 2),
  119.             dtype=jnp.float32,
  120.         )
  121.  
  122.         kernel = self.param(
  123.             "kernel",
  124.             nn.initializers.zeros,
  125.             (self.kernel_size, 1, 1),
  126.             PARAM_DTYPE,
  127.         )
  128.  
  129.         # For kernel size 3, this produces the mask [1, 0, 1].
  130.         kernel_mask = jnp.ones_like(kernel)
  131.         kernel_mask = kernel_mask.at[
  132.             self.kernel_size // 2, :, :
  133.         ].set(0)
  134.         kernel = kernel * kernel_mask
  135.  
  136.         logits = jax.lax.conv_general_dilated(
  137.             channel_mean.astype(COMPUTE_DTYPE)[..., None],
  138.             kernel.astype(COMPUTE_DTYPE),
  139.             window_strides=(1,),
  140.             padding="SAME",  # Zero padding along the channel axis.
  141.             dimension_numbers=("NWC", "WIO", "NWC"),
  142.         ).squeeze(-1)
  143.  
  144.         scale = (
  145.             2.0 * jax.nn.sigmoid(logits.astype(jnp.float32))
  146.         ).astype(COMPUTE_DTYPE)
  147.  
  148.         x = x * scale.reshape(batch, 1, 1, channels)
  149.         return x
  150.  
  151.  
  152. class PerChannelGate(nn.Module):
  153.     """One independent learned weight for each channel."""
  154.  
  155.     @nn.compact
  156.     def __call__(self, x):
  157.         batch, _height, _width, channels = x.shape
  158.  
  159.         channel_mean = jnp.mean(
  160.             x,
  161.             axis=(1, 2),
  162.             dtype=jnp.float32,
  163.         )
  164.  
  165.         channel_weights = self.param(
  166.             "channel_weights",
  167.             nn.initializers.zeros,
  168.             (channels,),
  169.             PARAM_DTYPE,
  170.         )
  171.  
  172.         logits = (
  173.             channel_mean *
  174.             channel_weights.astype(jnp.float32)
  175.         )
  176.  
  177.         scale = (
  178.             2.0 * jax.nn.sigmoid(logits)
  179.         ).astype(COMPUTE_DTYPE)
  180.  
  181.         x = x * scale.reshape(batch, 1, 1, channels)
  182.         return x
Advertisement
Add Comment
Please, Sign In to add comment