from d2l import jax as d2l
from flax import nnx
import jaxWe’ve seen a sequence of hand-designed architectures (LeNet → AlexNet → VGG → GoogLeNet → ResNet → DenseNet), each a hypothesis about what makes nets work.
Can we design networks more systematically?
RegNet (Radosavovic et al., 2020):
AnyNet): same template, free hyperparameters.Simple closed-form rules (“width grows linearly with stage”) outperform years of expert tuning.
The AnyNet design space.
Stem (low-level conv) → 4 stages of residual blocks → head (global pool + linear). Each stage’s depth, width, group count are free parameters:
The stem is deliberately plain: one stride-2 3×3 convolution, BatchNorm, ReLU. Its job is to halve resolution and create the first feature channels before the repeated stages begin.
class AnyNet(d2l.Classifier):
def __init__(self, arch, stem_channels, lr=0.1, num_classes=10,
in_channels=1, rngs=None):
super().__init__()
self.save_hyperparameters(ignore=['rngs'])
rngs = nnx.Rngs(d2l.get_key()) if rngs is None else rngs
self.net = self.create_net(in_channels, rngs)
def stem(self, in_channels, num_channels, rngs):
return nnx.Sequential(
nnx.Conv(in_channels, num_channels, kernel_size=(3, 3),
strides=(2, 2), padding=(1, 1), rngs=rngs),
nnx.BatchNorm(num_channels, rngs=rngs), nnx.relu)Each stage repeats the same ResNeXt block. The first block uses stride 2 and a 1×1 skip projection to change resolution and channel count; the rest preserve shape.
def stage(self, depth, num_channels, groups, bot_mul, in_channels, rngs):
blk = []
for i in range(depth):
if i == 0:
blk.append(d2l.ResNeXtBlock(num_channels, groups, bot_mul,
use_1x1conv=True, strides=(2, 2), in_channels=in_channels,
rngs=rngs))
else:
blk.append(d2l.ResNeXtBlock(num_channels, groups, bot_mul,
in_channels=num_channels, rngs=rngs))
return nnx.Sequential(*blk)The architecture tuple supplies (depth, channels, groups, bottleneck) per stage. The head is the now-standard global average pool + linear classifier.
def create_net(self, in_channels, rngs):
layers = [self.stem(in_channels, self.stem_channels, rngs)]
stage_channels = self.stem_channels
for s in self.arch:
layers.append(self.stage(*s, stage_channels, rngs))
stage_channels = s[1]
layers.append(nnx.Sequential(
lambda x: x.mean(axis=(1, 2)), # global avg pooling over H, W (NHWC)
nnx.Linear(stage_channels, self.num_classes, rngs=rngs)))
return nnx.Sequential(*layers)Comparing error empirical distribution functions of design spaces.
RegNet narrows AnyNet with simple constraints: stage widths grow approximately linearly, bottleneck ratios stay fixed, and group widths are shared across stages. The result is a smaller search space with better probability of good models.
The paper’s empirical findings collapse to: width grows linearly with stage, depth stays roughly constant, ResNeXt-style groups. A scaled-down version for Fashion-MNIST:
The architecture is competitive with hand-designed ResNets at similar parameter counts, and the discovery process scales trivially with compute.