Experiments

Repository workflows

Component Entry point Task
Decoder and DAMCHA models/transformer_damcha.py Causal sequence modeling
Vision Transformer models/vit.py Image classification model API
CV trainer train.py / main.py Masked patch reconstruction on CIFAR-10/100 and MNIST
NLP trainer train_nlp.py / main.py Source-conditioned WMT14 and CommonGen generation
Baselines models/baselines.py THA, DCMHA, ColMHA, MMA and MoA

Paper settings

Section 5.1 specifies the following decoder configuration:

Setting Manuscript
Decoder layers 10
Attention heads 8
Model / head dimension 512 / 64
Training epochs 50
Batch size 128
Optimizer Adam
Learning rate 0.0003
Metric MLP 9 layers
Text datasets WMT14 English–German, CommonGen
Image representation 1024 spatial tokens per 32 × 32 image, 8-bit pixels

The CLI provides the model-dimension, layer, batch-size and learning-rate controls. --mlp_hidden_dims lists the hidden widths; the final output projection adds one layer. Thus eight comma-separated hidden widths define a nine-layer generator. Widths and feed-forward size are experiment configuration choices; Section 5.1 specifies the generator depth.

The repository CV workflow uses masked patch reconstruction. Its default 4 × 4 patches produce 64 tokens per 32 × 32 image. The paper's autoregressive pixel-generation experiment is a separate task specification. The commands here run the existing reconstruction workflow. The current release covers the modules and experiment pipelines present in this repository.

Run comparisons

bash scripts/run_all_experiments.sh --dry-run
bash scripts/run_all_experiments.sh
SKIP_DONE=1 bash scripts/run_all_experiments.sh

The suite runs DAMCHA, standard MHA and five baselines on four datasets: 28 runs. It uses the CLI's compact configuration (D=128, four layers, two hidden MLP layers of widths 256 and 512). Set PYTHON, EPOCHS, CV_BATCH, NLP_BATCH, NUM_WORKERS, SEED or LOG_DIR as environment variables. Successful runs receive a .done marker; failed runs make the script exit with an error.

Evaluation

Images: pixel MSE, distribution FID from torchvision's ImageNet Inception-v3 features, and squared feature distances for individual image pairs. Feature-distance quartiles describe paired feature errors. Top-MSE subsets are ranked by reconstruction error, which corresponds to fixed-variance Gaussian reconstruction likelihood. Their FID is evaluated as a distribution.

Text: target-token NLL under source conditioning, greedy autoregressive generation, and mean log P(reference | hypothesis) from the configured BART scorer. BARTScore@Top-50 selects the 50 examples with lowest NLL, then averages their BARTScores. The scorer implementation follows the conditional log-probability convention of BARTScore; this repository's scorer checkpoint is facebook/bart-base.

CSV files retain per-epoch metrics; JSON files contain per-example metrics. Export each run's best validation epoch:

python -m scripts.parse_results --input outputs --output logs/experiments/results.csv

Profiling

python -m scripts.profile_flops_latency --device cuda --batch-size 8 --tokens 64
python -m scripts.profile_memory --device cuda:0 --batch-size 8 --tokens 64

Profiling measures the current decoder implementation with the CLI's compact architecture. FLOPs cover matrix-product and convolution operators recognized by torch.profiler; latency covers the whole forward pass. The CSV records device, batch size and sequence length. Memory profiling reports CUDA peak allocated MiB for inference and backward passes.

Regression checks

pytest -q

Tests cover the paper's block-row formula and gradients, input adaptation, prefix causality, shared-generator reuse, parameter registration, baseline masking, conditional generation, EOS handling and complete checkpoint restoration. The selected results on the documentation homepage are manuscript values; local regression checks validate implementation behavior.