Variants

UNet3D

20.8M

3D volumetric version using 3D convolutions, suitable for all 3D reservoir datasets (Arena, CO2 Nested, 3D Channels).

Conv: Conv3d
Norm: GroupNorm
Exact params: 20,827,592

UNet2D

7.3M

2D version for Cartesian grids, used on Two-Phase Oil-Water and CO2 Radial datasets.

Conv: Conv2d
Norm: GroupNorm
Exact params: 7,300,752

Training Provenance

These values are read from the generated benchmark result metadata, not from static method-page copy. When a model is retrained and the website data is regenerated, this section updates with the result records.

Arena training recipe

Epochs: 200
Batch size: 32
Learning rate: 0.0001
Loss: MSEGradientLoss, CombinedGradientLoss

Arena feature sets

Flow Arena input channels vary by fault complexity. Current result records for this method include 3 feature-set definitions.

PORO PERMX SWAT_0 PRESSURE_0 MULTX MULTY DEPTH

Strengths & Weaknesses

Strengths

  • Best-performing architecture on arena benchmarks and consistently strong across all dataset families
  • Skip connections preserve fine spatial details needed for sharp saturation fronts
  • Residual blocks with GroupNorm enable stable training even with small batch sizes
  • Relatively fast training convergence compared to transformer-based architectures
  • Well-understood architecture with extensive literature on tuning and modifications

Weaknesses

  • Purely local receptive field — may miss long-range pressure communication across the reservoir
  • Fixed spatial resolution hierarchy — cannot adapt to multi-scale features without architectural changes
  • Parameter count scales steeply with 3D convolutions (20.8M for 3D vs 7.3M for 2D)
  • No explicit mechanism for learning physical symmetries or conservation laws

Results Across Benchmarks

All results for U-Net variants across all benchmark families. Sorted by the first field's rel-L² error by default.

Flow Arena

Property
Fault
Control
Loss
28 results
    Sat (frac)Pres (bar) 
#DatasetModelLossEpochsrel-L²MREMAErel-L²MREMAEParams
1geostat_nf_rateUNet3DAbsLp(p=2)2000.14360.26280.0386770.03020.02605.74169720.8M
2geostat_nf_bhpUNet3DAbsLp(p=2)2000.14700.28330.0363870.02190.01844.65881120.8M
3geostat_nf_bhpUNet3DMSE2000.14780.27610.0359090.02280.01924.85158320.8M
4geostat_nf_rateUNet3DMSE2000.15550.30800.0392460.03130.02645.80587920.8M
5channels_nf_bhpUNet3DAbsLp(p=2)2000.17220.32360.0484370.02670.02205.43854820.8M
6channels_nf_bhpUNet3DMSE2000.17660.34650.0482610.02660.02175.39326720.8M
7geostat_zt_rateUNet3DAbsLp(p=2)2000.18480.33810.0469760.04560.03608.48619420.8M
8geostat_zt_rateUNet3DMSE2000.18540.34180.0461110.04690.03658.56901120.8M
9geostat_vt_rateUNet3DMSE2000.18940.35190.0544280.04610.03598.48650235.9M
10geostat_zt_bhpUNet3DMSE2000.18940.33680.0442860.03070.02486.21029420.8M
11geostat_zt_bhpUNet3DAbsLp(p=2)2000.19010.34730.0458460.03070.02516.27767220.8M
12geostat_vt_bhpUNet3DMSE2000.19420.36730.0530380.02950.02395.98858035.9M
13geostat_vt_rateUNet3DAbsLp(p=2)2000.20100.36290.0559840.04670.03648.45887135.9M
14geostat_vt_bhpUNet3DAbsLp(p=2)2000.20220.38080.0536070.03030.02436.07724735.9M
15channels_nf_rateUNet3DMSE2000.20890.38160.0527280.04370.03487.55857420.8M
16channels_nf_rateUNet3DAbsLp(p=2)2000.21040.38610.0545130.04490.03547.78299620.8M
17geostat_vt_rateUNet3DMSE+grad(w=1)2000.21090.37440.0545450.04810.03728.50827535.9M
18geostat_vt_bhpUNet3DMSE+grad(w=1)2000.21340.35900.0522820.03410.02736.82054435.9M
19channels_zt_bhpUNet3DMSE2000.22300.41640.0545510.03190.02496.21704520.8M
20channels_zt_bhpUNet3DAbsLp(p=2)2000.22320.43140.0559960.03230.02596.43961020.8M
21channels_vt_bhpUNet3DMSE2000.23840.45670.0657440.03440.02666.69766035.9M
22channels_vt_bhpUNet3DMSE+grad(w=1)2000.24020.47340.0671570.03440.02696.75383635.9M
23channels_zt_rateUNet3DAbsLp(p=2)2000.24570.45700.0622720.05370.04139.74435020.8M
24channels_zt_rateUNet3DMSE2000.24570.44660.0592860.05490.042410.07036620.8M
25channels_vt_bhpUNet3DAbsLp(p=2)2000.25060.52510.0699340.03540.02766.90482935.9M
26channels_vt_rateUNet3DMSE+grad(w=1)2000.25740.47500.0713650.05400.04099.88945135.9M
27channels_vt_rateUNet3DMSE2000.25960.50280.0735840.05520.041810.17899935.9M
28channels_vt_rateUNet3DAbsLp(p=2)2000.26650.54760.0761830.05640.043410.47856035.9M

Two-Phase Oil-Water

Loss
16 results
    Pres (bar)Sat (frac) 
#TargetModelLossEpochsrel-L²MREMAErel-L²MREMAEParams
1combinedUNet3DBadawiCombined20000.01030.00861.4302300.01820.01400.00476520.8M
2combinedUNet3DBadawiCombined2000.01420.01242.1764280.02540.02070.00715620.8M
3combinedUNet2DBadawiCombined40000.01430.01212.0387410.02020.01510.00513229.1M
4combinedUNet2DBadawiCombined40000.01500.01282.1364730.02000.01470.00503216.4M
5combinedUNet2DBadawiCombined40000.01650.01422.3641490.02090.01540.00524029.1M
6combinedUNet2DBadawiCombined5000.01790.01492.4694670.02390.01800.00617516.4M
7pressureUNet2DBadawiSingleField5000.02150.01923.150069---16.4M
8combinedUNet2DRelLp2000.02430.02043.4150110.02680.02010.0067987.3M
9pressureUNet3DBadawiSingleField2000.02840.02684.242195---20.8M
10combinedUNet2DRelLp2000.02840.02594.4944010.02430.01870.00645516.4M
11pressureUNet2DRelLp2000.03170.02784.753270---7.3M
12pressureUNet2DRelLp2000.03510.03165.364382---16.4M
13saturationUNet2DRelLp200---0.03190.02430.0081137.3M
14saturationUNet2DRelLp200---0.02990.02270.00762816.4M
15saturationUNet2DBadawiSingleField500---0.03020.02340.00788016.4M
16saturationUNet3DBadawiSingleField200---0.02400.01810.00626420.8M

CO2 Radial

Loss
23 results
    Pres (bar)Sat (frac) 
#TargetModelLossEpochsrel-L²MREMAErel-L²MREMAEParams
1pressureUNet3DUFNODerivLoss5000.00300.00080.181463---20.8M
2pressureUNet2DRelLp5000.00310.00080.193750---16.3M
3pressureUNet2DRelLp5000.00320.00100.217995---29.0M
4combinedUNet2DRelLp5000.00340.00100.2327920.07085.96540.00171316.3M
5pressureUNet2DRelLp5000.00370.00100.235549---29.2M
6pressureUNet2DUFNODerivLoss5000.00370.00110.248648---29.0M
7pressureUNet2DUFNODerivLoss5000.00390.00110.250622---16.3M
8pressureUNet2DUFNODerivLoss5000.00390.00110.246458---16.3M
9pressureUNet3DRelLp5000.00400.00130.286640---20.8M
10pressureUNet3DRelLp5000.00420.00140.291506---20.8M
11pressureUNet3DUFNODerivLoss5000.00450.00140.296649---20.8M
12saturationUNet2DRelLp500---0.06536.65510.00172616.3M
13saturationUNet2DRelLp500---0.08118.28270.00247129.2M
14saturationUNet2DRelLp500---0.08118.28270.00247129.2M
15saturationUNet2DUFNODerivLoss500---0.08296.28470.00218216.3M
16saturationUNet2DUFNODerivLoss500---0.08046.48450.00214116.3M
17saturationUNet2DRelLp500---0.07035.27980.00167529.0M
18saturationUNet2DRelLp500---0.07035.27980.00167529.0M
19saturationUNet2DUFNODerivLoss500---0.07056.54780.00176729.0M
20saturationUNet3DRelLp500---0.06576.11840.00171520.8M
21saturationUNet3DRelLp500---0.10055.66920.00166720.8M
22saturationUNet3DUFNODerivLoss500---0.09546.23120.00208320.8M
23saturationUNet3DUFNODerivLoss500---0.06344.01320.00115120.8M

CO2 Nested

3 results
   Pres (bar) 
#ModelLossEpochsrel-L²MREMAEParams
1UNet3DRelLp8004.12e-58.46e-60.00186420.8M
2UNet3DRelLp5004.25e-58.99e-60.00198220.8M
3UNet3DRelLp2005.55e-51.40e-50.00308320.8M

Two-Phase 3D Channels

6 results
    Pres (bar)Sat (frac) 
#TargetModelLossEpochsrel-L²MREMAErel-L²MREMAEParams
1pressureUNet3DRelLp5000.00070.00050.163858---20.8M
2pressureUNet3DRelLp2000.00070.00060.181097---20.8M
3combinedUNet3DRelLp5000.00090.00060.2067610.12780.07580.01582520.8M
4combinedUNet3DRelLp2000.00100.00070.2311140.13200.08990.01801520.8M
5saturationUNet3DRelLp200---0.13230.08070.01717920.8M
6saturationUNet3DRelLp500---0.12980.07090.01552820.8M

Training Curves

UNet3D - Flow Arena (channels_nf_bhp)

UNet3D - Flow Arena (channels_nf_rate)

UNet3D - Flow Arena (channels_vt_bhp)

UNet3D - Flow Arena (channels_vt_rate)

UNet3D - Flow Arena (channels_zt_bhp)

UNet3D - Flow Arena (channels_zt_rate)

UNet3D - Flow Arena (geostat_nf_bhp)

UNet3D - Flow Arena (geostat_nf_rate)

Architecture

U-Net was originally introduced by Ronneberger et al. (2015) for biomedical image segmentation and has since become the de facto standard architecture for dense prediction tasks. Its encoder-decoder structure with skip connections allows it to combine high-level semantic features from the bottleneck with fine-grained spatial details from the encoder — critical for accurately predicting spatially varying fields like pressure and saturation.

Our implementation uses residual blocks (ResNet-style) with GroupNorm normalization instead of BatchNorm, which provides more stable training with small batch sizes common in 3D volumetric applications. The architecture follows a standard 4-level hierarchy with channel doubling at each downsampling stage.

For 3D datasets, we use volumetric 3D convolutions throughout (UNet3D). For 2D datasets, we use standard 2D convolutions (UNet2D). Both variants share the same architectural design — residual blocks, skip connections, and GroupNorm — differing only in the spatial dimensionality of the convolution kernels.

U-Net serves as our primary baseline architecture. Its consistent performance across all benchmarks makes it the reference point against which more specialized architectures are measured.

Model Summary

UNet3D

PaddedUNet3D(
  (conv_in): Conv3d(5, 32, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
  (down_blocks): ModuleList(
    (0): DownBlock3D(
      (resnets): ModuleList(
        (0): ResnetBlock3D(
          (conv1): Conv3d(32, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv3d(64, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Conv3d(32, 64, kernel_size=(1, 1, 1), stride=(1, 1, 1))
        )
        (1): ResnetBlock3D(
          (conv1): Conv3d(64, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv3d(64, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (downsample): Conv3d(64, 64, kernel_size=(2, 2, 2), stride=(2, 2, 2))
    )
    (1): DownBlock3D(
      (resnets): ModuleList(
        (0): ResnetBlock3D(
          (conv1): Conv3d(64, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv3d(128, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Conv3d(64, 128, kernel_size=(1, 1, 1), stride=(1, 1, 1))
        )
        (1): ResnetBlock3D(
          (conv1): Conv3d(128, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv3d(128, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (downsample): Conv3d(128, 128, kernel_size=(2, 2, 2), stride=(2, 2, 2))
    )
    (2): DownBlock3D(
      (resnets): ModuleList(
        (0): ResnetBlock3D(
          (conv1): Conv3d(128, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
          (conv2): Conv3d(256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
          (shortcut): Conv3d(128, 256, kernel_size=(1, 1, 1), stride=(1, 1, 1))
        )
        (1): ResnetBlock3D(
          (conv1): Conv3d(256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
          (conv2): Conv3d(256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (downsample): Conv3d(256, 256, kernel_size=(2, 2, 2), stride=(2, 2, 2))
    )
  )
  (mid_resnet1): ResnetBlock3D(
    (conv1): Conv3d(256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
    (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
    (conv2): Conv3d(256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
    (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
    (shortcut): Identity()
  )
  (mid_resnet2): ResnetBlock3D(
    (conv1): Conv3d(256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
    (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
    (conv2): Conv3d(256, 256, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
    (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
    (shortcut): Identity()
  )
  (up_blocks): ModuleList(
    (0): UpBlock3D(
      (resnets): ModuleList(
        (0): ResnetBlock3D(
          (conv1): Conv3d(512, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv3d(128, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Conv3d(512, 128, kernel_size=(1, 1, 1), stride=(1, 1, 1))
        )
        (1): ResnetBlock3D(
          (conv1): Conv3d(128, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv3d(128, 128, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (upsample): ConvTranspose3d(256, 256, kernel_size=(2, 2, 2), stride=(2, 2, 2))
    )
    (1): UpBlock3D(
      (resnets): ModuleList(
        (0): ResnetBlock3D(
          (conv1): Conv3d(256, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv3d(64, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Conv3d(256, 64, kernel_size=(1, 1, 1), stride=(1, 1, 1))
        )
        (1): ResnetBlock3D(
          (conv1): Conv3d(64, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv3d(64, 64, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (upsample): ConvTranspose3d(128, 128, kernel_size=(2, 2, 2), stride=(2, 2, 2))
    )
    (2): UpBlock3D(
      (resnets): ModuleList(
        (0): ResnetBlock3D(
          (conv1): Conv3d(128, 32, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(1, 32, eps=1e-05, affine=True)
          (conv2): Conv3d(32, 32, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(1, 32, eps=1e-05, affine=True)
          (shortcut): Conv3d(128, 32, kernel_size=(1, 1, 1), stride=(1, 1, 1))
        )
        (1): ResnetBlock3D(
          (conv1): Conv3d(32, 32, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm1): GroupNorm(1, 32, eps=1e-05, affine=True)
          (conv2): Conv3d(32, 32, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
          (norm2): GroupNorm(1, 32, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (upsample): ConvTranspose3d(64, 64, kernel_size=(2, 2, 2), stride=(2, 2, 2))
    )
  )
  (conv_out): Conv3d(32, 40, kernel_size=(3, 3, 3), stride=(1, 1, 1), padding=(1, 1, 1))
)

UNet2D

PaddedUNet2D(
  (conv_in): Conv2d(84, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
  (down_blocks): ModuleList(
    (0): DownBlock2D(
      (resnets): ModuleList(
        (0): ResnetBlock2D(
          (conv1): Conv2d(32, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Conv2d(32, 64, kernel_size=(1, 1), stride=(1, 1))
        )
        (1): ResnetBlock2D(
          (conv1): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (downsample): Conv2d(64, 64, kernel_size=(2, 2), stride=(2, 2))
    )
    (1): DownBlock2D(
      (resnets): ModuleList(
        (0): ResnetBlock2D(
          (conv1): Conv2d(64, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Conv2d(64, 128, kernel_size=(1, 1), stride=(1, 1))
        )
        (1): ResnetBlock2D(
          (conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (downsample): Conv2d(128, 128, kernel_size=(2, 2), stride=(2, 2))
    )
    (2): DownBlock2D(
      (resnets): ModuleList(
        (0): ResnetBlock2D(
          (conv1): Conv2d(128, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
          (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
          (shortcut): Conv2d(128, 256, kernel_size=(1, 1), stride=(1, 1))
        )
        (1): ResnetBlock2D(
          (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
          (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (downsample): Conv2d(256, 256, kernel_size=(2, 2), stride=(2, 2))
    )
  )
  (mid_resnet1): ResnetBlock2D(
    (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
    (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
    (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
    (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
    (shortcut): Identity()
  )
  (mid_resnet2): ResnetBlock2D(
    (conv1): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
    (norm1): GroupNorm(8, 256, eps=1e-05, affine=True)
    (conv2): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
    (norm2): GroupNorm(8, 256, eps=1e-05, affine=True)
    (shortcut): Identity()
  )
  (up_blocks): ModuleList(
    (0): UpBlock2D(
      (resnets): ModuleList(
        (0): ResnetBlock2D(
          (conv1): Conv2d(512, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Conv2d(512, 128, kernel_size=(1, 1), stride=(1, 1))
        )
        (1): ResnetBlock2D(
          (conv1): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(4, 128, eps=1e-05, affine=True)
          (conv2): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(4, 128, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (upsample): ConvTranspose2d(256, 256, kernel_size=(2, 2), stride=(2, 2))
    )
    (1): UpBlock2D(
      (resnets): ModuleList(
        (0): ResnetBlock2D(
          (conv1): Conv2d(256, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Conv2d(256, 64, kernel_size=(1, 1), stride=(1, 1))
        )
        (1): ResnetBlock2D(
          (conv1): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(2, 64, eps=1e-05, affine=True)
          (conv2): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(2, 64, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (upsample): ConvTranspose2d(128, 128, kernel_size=(2, 2), stride=(2, 2))
    )
    (2): UpBlock2D(
      (resnets): ModuleList(
        (0): ResnetBlock2D(
          (conv1): Conv2d(128, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(1, 32, eps=1e-05, affine=True)
          (conv2): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(1, 32, eps=1e-05, affine=True)
          (shortcut): Conv2d(128, 32, kernel_size=(1, 1), stride=(1, 1))
        )
        (1): ResnetBlock2D(
          (conv1): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm1): GroupNorm(1, 32, eps=1e-05, affine=True)
          (conv2): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
          (norm2): GroupNorm(1, 32, eps=1e-05, affine=True)
          (shortcut): Identity()
        )
      )
      (upsample): ConvTranspose2d(64, 64, kernel_size=(2, 2), stride=(2, 2))
    )
  )
  (conv_out): Conv2d(32, 80, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
)