DifFRACT: Diffusion Feature Reconstruction and Attribution for Circuit Tracing

Trained timestep-conditioned transcoders for three multimodal diffusion transformers (MM-DiT): FLUX.1 [schnell], FLUX.1 [dev] and Stable Diffusion 3.5 Medium, plus the SAE baselines for FLUX.1 [schnell]. They accompany the paper DifFRACT: Diffusion Feature Reconstruction and Attribution for Circuit Tracing (arXiv:2606.15796).

A transcoder decomposes an MLP sublayer into a sparse linear combination of interpretable features; conditioning it on the denoising timestep lets one transcoder track how a feature behaves across the whole diffusion trajectory. Substituting the transcoders into a frozen Local Replacement Model yields the attribution graphs and the circuit-guided interventions studied in the paper. The code that loads and uses these weights lives in the companion repository: github.com/Artalmaz31/DifFRACT.

Contents

One folder of transcoders per backbone, plus the SAE baselines for FLUX.1 [schnell]. Every file is a state_dict for a TemporalAwareTranscoder module (the SAE baseline shares the identical architecture).

Folder Backbone Files Layers
flux-schnell-transcoders/ FLUX.1 [schnell] 34 0-15 and 18, both streams
flux-schnell-saes/ FLUX.1 [schnell] 6 6, 12, 18, both streams
flux-dev-transcoders/ FLUX.1 [dev] 32 0-15, both streams
sd3-5-medium-transcoders/ SD 3.5 Medium 32 0-15, both streams

The 32 transcoders for layers 0-15 of each backbone are the set the Local Replacement Model runs with (the paper's case studies use the FLUX.1 [schnell] set). FLUX.1 [schnell] layer 18 and the SAEs at layers 6 / 12 / 18 support the sparsity-faithfulness comparison.

hf download Artalmaz31/DifFRACT --include "flux-schnell-transcoders/*" --local-dir weights

Model architecture

Each module maps an MLP input xx to its output y^\hat{y}, conditioned on the diffusion timestep tt:

  • a sinusoidal embedding of tt, a 2-layer SiLU MLP, then a linear head producing the FiLM pair (scale,shift)(\mathrm{scale}, \mathrm{shift});
  • modulation xmod=xβŠ™(1+scale)+shiftx_{\mathrm{mod}} = x \odot (1 + \mathrm{scale}) + \mathrm{shift};
  • a ReLU encoder z=ReLU(Wenc xmod+benc)z = \mathrm{ReLU}(W_{\mathrm{enc}}\,x_{\mathrm{mod}} + b_{\mathrm{enc}}), the sparse code;
  • a linear decoder with unit-norm columns y^=Wdec z+bdec\hat{y} = W_{\mathrm{dec}}\,z + b_{\mathrm{dec}}.

The SAE baseline is architecturally identical but autoencodes the MLP output (input = target), so its reconstruction error is directly comparable to a transcoder's.

Training recipe

FLUX.1 [schnell] FLUX.1 [dev] SD 3.5 Medium
Residual width d_model 3072 3072 1536
Expansion factor / d_feat 16 / 49152 16 / 49152 16 / 24576
Timestep embedding dim 256 256 256
Sparsity (L1) Ξ»img\lambda_{\mathrm{img}} / Ξ»txt\lambda_{\mathrm{txt}} 3e-4 / 5e-5 3e-4 / 5e-5 3e-4 / 5e-5
Reconstruction loss variance-normalized MSE variance-normalized MSE variance-normalized MSE
Optimizer AdamW, lr 2e-4, wd 0, cosine AdamW, lr 2e-4, wd 0, cosine AdamW, lr 2e-4, wd 0, cosine
Steps / guidance / resolution 4 / 0 / 512Γ—512512 \times 512 50 / 3.5 / 512Γ—512512 \times 512 40 / 4.5 / 512Γ—512512 \times 512
Prompts yvdao/midjourney-v6 yvdao/midjourney-v6 yvdao/midjourney-v6

Usage

Install the companion code, then:

from huggingface_hub import snapshot_download
from transcoder_training.transcoder import load_transcoders

path = snapshot_download("Artalmaz31/DifFRACT", allow_patterns=["flux-schnell-transcoders/*"])
transcoders = load_transcoders(
    f"{path}/flux-schnell-transcoders",
    layers=range(16),
    d_model=3072,
    expansion_factor=16,
    time_embed_dim=256,
)

The Local Replacement Model, attribution graphs and interventions take the backbone name and the folder:

from transcoder_circuits.replacement_model import LRMConfig
from transcoder_circuits.circuit_analysis import LRMPipeline

cfg = LRMConfig.for_model(
    "flux-schnell",  # or "flux-dev", "sd3.5-medium"
    transcoder_dir=f"{path}/flux-schnell-transcoders",
    target_layers=tuple(range(16)),
)
pipeline = LRMPipeline(cfg)
pipeline.initialize()
pipeline.load_transcoders()

An individual SAE baseline:

import torch
from transcoder_training.transcoder import TemporalAwareSAE

sae = TemporalAwareSAE(d_model=3072, expansion_factor=16, time_embed_dim=256)
sae.load_state_dict(torch.load(f"{path}/flux-schnell-saes/sae_img_12.pt", map_location="cpu"))
sae.eval()

walkthrough.ipynb in the companion repository runs the end-to-end pipeline on FLUX.1 [schnell]: Local Replacement Model, attribution graph, pruning, interactive visualization and a circuit-guided intervention.

Citation

@misc{mazur2026diffractdiffusionfeaturereconstruction,
      title={DifFRACT: Diffusion Feature Reconstruction and Attribution for Circuit Tracing},
      author={Artyom Mazur and Nina Konovalova and Aibek Alanov},
      year={2026},
      eprint={2606.15796},
      archivePrefix={arXiv},
      primaryClass={cs.CV},
      url={https://arxiv.org/abs/2606.15796},
}
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 Artalmaz31/DifFRACT

Finetuned
(610)
this model

Dataset used to train Artalmaz31/DifFRACT

Paper for Artalmaz31/DifFRACT