MCPcopy Create free account
hub / github.com/ml-explore/mlx / BenchNetMLX

Class BenchNetMLX

benchmarks/python/conv2d_train_bench_cpu.py:12–35  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

10 mx.set_default_device(mx.cpu)
11
12 class BenchNetMLX(mlx.nn.Module):
13 # simple encoder-decoder net
14
15 def __init__(self, in_channels, hidden_channels=32):
16 super().__init__()
17
18 self.net = mlx.nn.Sequential(
19 mlx.nn.Conv2d(in_channels, hidden_channels, kernel_size=3, padding=1),
20 mlx.nn.ReLU(),
21 mlx.nn.Conv2d(
22 hidden_channels, 2 * hidden_channels, kernel_size=3, padding=1
23 ),
24 mlx.nn.ReLU(),
25 mlx.nn.ConvTranspose2d(
26 2 * hidden_channels, hidden_channels, kernel_size=3, padding=1
27 ),
28 mlx.nn.ReLU(),
29 mlx.nn.ConvTranspose2d(
30 hidden_channels, in_channels, kernel_size=3, padding=1
31 ),
32 )
33
34 def __call__(self, input):
35 return self.net(input)
36
37 benchNet = BenchNetMLX(3)
38 mx.eval(benchNet.parameters())

Callers 1

bench_mlxFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected