Skip to content

AChenreddy24/RobustQuantization

Folders and files

NameName
Last commit message
Last commit date

Latest commit

 

History

1 Commit
 
 
 
 
 
 
 
 
 
 

Repository files navigation

Quantization Bias Testbed — ResNet + DDPM

Minimal, self-contained code for testing FP baselines and post-training-quantized variants of two model families:

  1. CIFAR-100 ResNet-20 classifiers — FP baseline, naive uniform quantization, and Class-Weighted Output Recovery (CWOR) calibration for W4/W2 activation quant.
  2. CIFAR-10 DDPM (google/ddpm-cifar10-32) — FP FID baseline, per-tensor and per-channel quantization, BRECQ/AdaRound layer reconstruction, and per-class quantization-damage diagnostics.

Layout

publishable/
├── src/                     # reusable library
│   ├── data/cv_data.py           # CIFAR-100/10 loaders + corruption transforms
│   ├── quantization/
│   │   ├── quantizer.py           # UniformQuantizer, QuantizedLayer
│   │   ├── wrap_generic.py        # wrap_model_with_quantizers()
│   │   └── reconstruction.py      # layer/block reconstruction helpers
│   ├── dro/objectives.py         # DRO losses (Wasserstein, KL, CVaR, χ²)
│   └── evaluation/metrics.py     # accuracy, per-class accuracy, CV robustness
└── scripts/                 # experiment runners
    ├── run_cv_experiment.py           # ResNet FP + naive-quant baseline (main entry)
    ├── run_cwor.py                    # CWOR calibration experiment
    ├── build_cwor_cali.py             # build class-weighted CIFAR calibration set
    ├── run_class_reweight_test.py     # per-class reweighting ablation
    ├── run_tail_class_diagnostic.py   # tail-class accuracy drop analysis
    ├── run_diffusion_cwor.py          # DDPM loading + FID + CWOR wrapper (core)
    ├── gen_qdiff_cali_cifar10.py      # generate DDIM-sampled calibration data
    ├── run_ddpm_fid_qdiff_baseline.py # naive FID sweep on DDPM under W{4,8}A{4,8}
    ├── run_ddpm_ptq_recon.py          # BRECQ/AdaRound layer reconstruction pipeline
    ├── run_ddpm_ptq_wa_recon.py       # weight + activation reconstruction variant
    ├── run_ddpm_fid_eval.py           # standalone FID eval for a saved ckpt
    ├── run_ddpm_per_class_diagnostic.py  # per-class damage on quantized DDPM
    ├── analyze_ddpm_per_class_bias.py    # aggregate per-class analysis
    ├── run_ddpm_qdiff_verify_and_visualize.py  # sanity + sample grids
    └── run_ddpm_visual_grid.py           # side-by-side FP vs quant sample grid

Setup

python -m venv .venv && source .venv/bin/activate
pip install -r requirements.txt

# For any script that uses Q-Diffusion's quantizers, get the Q-Diffusion repo
# and point QDIFFUSION_PATH at it:
git clone https://github.com/Xiuyu-Li/q-diffusion.git ~/q-diffusion
export QDIFFUSION_PATH=~/q-diffusion

Scripts add sys.path.insert(0, os.environ.get('QDIFFUSION_PATH', '~/q-diffusion')) so setting the env var (or dropping Q-Diffusion at ~/q-diffusion) is enough.

Quick-start — ResNet CIFAR-100

# FP + naive W4A8 quantization baseline
python scripts/run_cv_experiment.py --model resnet20 --w_bits 4 --a_bits 8

# CWOR class-weighted calibration
python scripts/build_cwor_cali.py --calib_size 512 --tail_weight 4.0 \
    --out results/cwor_cali.pt
python scripts/run_cwor.py --cali_path results/cwor_cali.pt --w_bits 4 --a_bits 8

Quick-start — CIFAR-10 DDPM

# 1. Generate FP DDIM calibration trajectory (requires Q-Diffusion)
python scripts/gen_qdiff_cali_cifar10.py \
    --config $QDIFFUSION_PATH/configs/cifar10.yml \
    --n_samples 1024 --timesteps 100 \
    --out results/qdiff_cifar10_cali.pt

# 2. Naive quantization FID sweep (per-tensor and per-channel)
python scripts/run_ddpm_fid_qdiff_baseline.py \
    --n_generate 5000 --ddim_steps 50 \
    --out_path results/ddpm_fid_baseline.json

# 3. BRECQ/AdaRound reconstruction
python scripts/run_ddpm_ptq_recon.py \
    --cali_path results/qdiff_cifar10_cali.pt \
    --w_bits 4 --a_bits 8 \
    --iters 2000 --out_dir results/ddpm_ptq_recon/

# 4. Per-class quantization damage diagnostic
python scripts/run_ddpm_per_class_diagnostic.py --w_bits 4 --a_bits 8
python scripts/analyze_ddpm_per_class_bias.py \
    --in_dir results/ddpm_per_class/ \
    --out_dir results/ddpm_per_class_analysis/

# 5. Side-by-side FP vs quantized visualization
python scripts/run_ddpm_visual_grid.py --n_prompts 12 --out results/ddpm_grid.png

Dependencies (see requirements.txt)

  • Core: torch, torchvision, numpy, pillow
  • Diffusion: diffusers, transformers, accelerate
  • FID / metrics: clean-fid, scipy
  • External (checkout separately, point via QDIFFUSION_PATH): Q-Diffusion — used only for its qdiff.quant_layer.UniformAffineQuantizer / AdaRoundQuantizer in the DDPM reconstruction scripts.

Reproducibility notes

  • All CIFAR ResNet experiments seed with --seed 42 by default.
  • DDPM FID uses clean-fid against 50k CIFAR-10 train images at 32×32.
  • DDPM PTQ reconstruction is layer-wise (BRECQ block wrappers not implemented for the diffusers UNet2DModel), so reported FIDs are typically 3–8 higher than Q-Diffusion's paper numbers for CIFAR-10 (which use their custom Model class with block reconstruction and split-shortcut). Details in docs/deltas.md if you add one — currently a design note.

About

No description, website, or topics provided.

Resources

Stars

Watchers

Forks

Releases

Packages

Contributors

Languages