import torch import torch.nn as nn # Deep network: 7 hidden layers, all tanh class DeepTanhNet(nn.Module): def __init__(self, dim, depth): super().__init__() layers = [] for _ in range(depth): layers.append(nn.Linear(dim, dim)) layers.append(nn.Tanh()) self.seq = nn.Sequential(*layers) self.out = nn.Linear(dim, 1) def forward(self, x): return self.out(self.seq(x)) depth = 7 dim = 8 net = DeepTanhNet(dim, depth) x = torch.randn(1, dim) target = torch.tensor([[0.0]]) out = net(x) loss = (out - target).pow(2).mean() loss.backward() # Print average absolute gradient for each Linear layer for i, layer in enumerate(net.seq): if isinstance(layer, nn.Linear): grad = layer.weight.grad.abs().mean().item() print(f"Layer {i//2 + 1} mean abs gradient: {grad:.6f}")