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.
- Code: https://github.com/mahishguru/microstructure-encoder-decoder
- Inverse design with these models (MERIDIAN): https://github.com/mahishguru/meridian
- Orientation codec (RGB ↔ HCP orientation field): https://github.com/mahishguru/orientation-codec
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.
Model tree for mahishguru/microstructure-encoder-decoder
Base model
laion/CLIP-ViT-H-14-laion2B-s32B-b79K