Microstructure encoder–decoder (ViT-H/14 + FM-DiT)

Trained weights for the generative encoder–decoder that learns a searchable latent space of grain morphology and HCP crystallographic texture in extruded magnesium alloys. A microstructure is represented as an RGB orientation image made by the orientation codec. The encoder compresses the image into a spatial latent z ∈ R^(16 × d_z), and the decoder turns any point of that latent space back into an orientation image. The image can then be converted into a DAMASK-ready RVE.

Contents

Path Model Latent z Notes
fmdit_512/fmdit_512.pth FM-DiT-512 16 × 32 = 512 used for the MERIDIAN design runs
fmdit_768/fmdit_768.pth FM-DiT-768 16 × 48 = 768
fmdit_1024/fmdit_1024.pth FM-DiT-1024 16 × 64 = 1,024
fmdit_1280/fmdit_1280.pth FM-DiT-1280 16 × 80 = 1,280
texture_priors/texture_prior_<variant>/{head.pt,codebook.npy} z → ODF-histogram head — used by the z-only ODF calibration step
texture_priors/texture_prior/{head.pt,codebook.npy} shared default texture prior — default codebook of the calibration step
baselines/vitdit.pth ViT-H/14 + DiT-XL/2 decoder 512 (global) baseline
baselines/vitsdxl.pth ViT-H/14 + frozen SDXL UNet 512 (global) baseline (trained adapters only)
baselines/vitvqgan.pth ViT-VQGAN none (discrete tokens) reference without a continuous bottleneck

The FM-DiT files contain only the trained weights: the full encoder (ViT-H/14, attention pooler, projection), the token adapter, the pooled projection, the texture head and the rank-64 LoRA matrices. The frozen Stable Diffusion 3.5 medium backbone and VAE are not redistributed. They are downloaded from stabilityai/stable-diffusion-3.5-medium when the decoder is built, so you must accept that model's licence on Hugging Face and set HF_TOKEN.

Architecture (FM-DiT)

  • Encoder: OpenCLIP ViT-H/14 (LAION-2B). The lower 16 blocks are frozen and the upper 16 fine-tuned. A 16-query attention pooler feeds a per-token projection to width d_z. Input is the 300 × 300 orientation image, upsampled with nearest-neighbour interpolation to 512 × 512.

  • Decoder: frozen SD3.5-medium MMDiT (2.5 B parameters), conditioned through:

    • a 5-block token adapter (16 tokens → joint-attention stream);
    • a pooled projection;
    • rank-64 LoRA on the attention and feed-forward projections, enabled from epoch 4.

    Rectified-flow velocity objective, plus reconstruction, spectral, InfoNCE, VICReg, latent-statistics, HCP orientation-distribution and texture-head terms.

  • Training: 40 epochs, AdamW, effective batch 576 on 4 × H200, CFG dropout 0.1.

  • Sampling (as evaluated): 25 Euler steps, shift 1.0, noise temperature 0.9, guidance 3.0.

Usage

git clone https://github.com/mahishguru/microstructure-encoder-decoder.git
cd microstructure-encoder-decoder && pip install -e ".[eval]"
export HF_TOKEN=...            # with access to stabilityai/stable-diffusion-3.5-medium
import os, torch
os.environ.update(FMDIT_LORA_RANK_OVERRIDE="64", FMDIT_LORA_ALPHA_OVERRIDE="64",
                  FMDIT_LORA_FFN="1", FMDIT_TEXHEAD="1")      # FM-DiT v4 architecture
from microstructure_ed.checkpoints import download, load_encoder_decoder
from microstructure_ed.encoder_arch_pretrained import Compressor
from microstructure_ed.fmdit.decoder_arch_pretrained import FlowMatchingDiTDecoder, SD35VAE

W = 512
encoder = Compressor(use_gradient_checkpointing=False, trainable_blocks=0, target_dim=W, spatial_tokens=16)
decoder = FlowMatchingDiTDecoder(target_dim=W)
load_encoder_decoder(encoder, decoder, download("vitfmdit"))  # "vitfmdit_768", "vitfmdit_1024", "vitfmdit_1280"
vae = SD35VAE(token=os.environ["HF_TOKEN"])
encoder, decoder, vae = encoder.cuda().eval(), decoder.cuda().eval(), vae.to("cuda")

x = ...                                   # (B, 3, 512, 512) orientation image in [-1, 1]
with torch.no_grad(), torch.amp.autocast("cuda", dtype=torch.bfloat16):
    z = encoder(x)                        # (B, 512)
    lat = decoder.sample(z, num_steps=25, noise_temp=0.9, shift=1.0, guidance_scale=3.0)
    img = vae.decode(lat)                 # (B, 3, 512, 512) in [-1, 1]

Resize the output to 300 × 300 with nearest-neighbour interpolation, then convert it into a DREAM.3D RVE with the orientation codec. The code repository covers the full evaluation and calibration pipeline.

Reconstruction fidelity

Held-out test set (n = 10,100 RVEs), at the 300 × 300 resolution used by the crystal-plasticity oracle. Grains are segmented identically on originals and reconstructions (Segment Anything, ViT-B). Values are mean ± standard deviation over samples; FID is a single distribution-level value.

Model FID ↓ MS-SSIM ↑ LPIPS ↓ Grain size rel. error ↓ Orientation EMD (°) ↓ Mean disorientation (°) ↓
FM-DiT-512 26.63 0.065 ± 0.040 0.576 ± 0.023 0.131 ± 0.120 8.606 ± 1.372 4.057 ± 2.086
FM-DiT-768 27.24 0.065 ± 0.038 0.575 ± 0.024 0.122 ± 0.115 8.515 ± 1.383 3.993 ± 2.041
FM-DiT-1024 25.37 0.064 ± 0.037 0.573 ± 0.023 0.101 ± 0.083 8.451 ± 1.329 3.938 ± 2.025
FM-DiT-1280 25.00 0.065 ± 0.039 0.574 ± 0.024 0.112 ± 0.107 8.398 ± 1.346 3.907 ± 1.973

Intended use and limitations

  • Intended use: research on generative modelling and inverse design of polycrystalline microstructure and texture, in particular latent-space optimisation against a crystal-plasticity oracle.
  • Statistical, not literal, reconstruction: decodes reproduce grain-size, grain-shape and texture statistics of the input, not the position of individual grains. Per-pixel image metrics therefore look weak.
  • Domain: 2D RVEs of extruded, single-phase HCP Mg alloys encoded with the bundled class-mean quaternions (class_means.json). Other materials, crystal systems or encodings are out of distribution.
  • Texture sharpness: the basal fibre is reproduced in the right place but somewhat under-sharpened. The texture-prior calibration step partly compensates for this.

Licences

The weights in this repository build on third-party models, and each file inherits the terms of the model it was fine-tuned from:

Files Terms
fmdit_*/*.pth (encoder, adapters, LoRA on SD3.5) Stability AI Community License for the parts derived from Stable Diffusion 3.5 (Powered by Stability AI); OpenCLIP ViT-H/14 encoder weights: MIT
baselines/vitdit.pth (fine-tuned DiT-XL/2) CC BY-NC 4.0, as the DiT-XL/2 weights (facebook/DiT-XL-2-256); non-commercial use only
baselines/vitsdxl.pth (adapters for SDXL) MIT for the released adapter weights; use with SDXL-base-1.0 is subject to the CreativeML Open RAIL++-M License
baselines/vitvqgan.pth (fine-tuned PaintMind ViT-VQGAN) Apache-2.0, as PaintMind
texture_priors/* MIT

Acknowledgements

Developed at the Institute of Material and Process Design, Helmholtz-Zentrum Hereon, Geesthacht, Germany. The ViT-VQGAN baseline builds on PaintMind by Qiyuan Ge.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support

Model tree for mahishguru/microstructure-encoder-decoder

Adapter
(1)
this model