This project compares beamforming methods for multi-user MIMO with fixed K_model = 8 user slots:
| Track | Methods | Entry scripts |
|---|---|---|
| MLP beamformers | WMMSE, DNN, BNN | scripts/train_dnn_bnn.py, scripts/evaluate_dnn_bnn.py |
| Deep unfolding | Truncated WMMSE, deep unfolding, Nesterov unfolding, Bayesian Nesterov | scripts/train_unfolded.py, scripts/evaluate_unfolded.py |
BayesianNeuralNetwork/
bnn_wmmse/ # importable Python package
config.py # shared hyperparameters
channel/data.py # channel generation, dataset loader
metrics/ # WSR (PyTorch + NumPy)
algorithms/wmmse.py # classical iterative WMMSE (NumPy)
models/dnn_bnn.py # DNN and BNN MLP beamformers
unfolding/ # deep-unfolded WMMSE family
pgd.py # shared PGD step
deep_unfolded.py
nesterov_unfolded.py
bayesian_nesterov_unfolded.py
factory.py
training/losses.py # supervised beamforming losses
plotting/ # plot helpers
utils/config_loader.py
scripts/ # CLI entry points
train_dnn_bnn.py
train_unfolded.py
evaluate_dnn_bnn.py
evaluate_unfolded.py
generate_dataset.py
plot_dnn_bnn.py
plot_unfolded.py
datasets/ # generated .npz (gitignored)
results_*/ # experiment outputs (gitignored)
python -m venv .venv
.\.venv\Scripts\Activate.ps1
pip install --upgrade pip
pip install -r requirements.txt
pip install -e .python scripts/generate_dataset.py --out_dir datasets --name wmmse_k4_debug --num_samples 2000 --batch_size 64 --k_active 4Supervised DNN:
python scripts/train_dnn_bnn.py --model dnn --objective supervised --dataset_path datasets/wmmse_k4_debug.npz --steps 2000 --out_dir debug_run --k_train 4Unsupervised DNN (fine-tune from checkpoint):
python scripts/train_dnn_bnn.py --model dnn --objective unsupervised --steps 50000 --batch_size 128 --out_dir results_k4_unsup_ft --k_train 4 --init_ckpt results_k4_50k/dnn/model.ptUnsupervised BNN:
python scripts/train_dnn_bnn.py --model bnn --objective unsupervised --steps 50000 --batch_size 128 --out_dir results_k4_unsup_ft --k_train 4 --kl_beta 1e-7 --init_ckpt results_k4_50k/dnn/model.ptActive users sweep:
python scripts/evaluate_dnn_bnn.py --run_dir results_k4_unsup_ft --bnn_samples 10 --k_test_min 1 --k_test_max 8SNR sweep (0–40 dB):
python scripts/evaluate_dnn_bnn.py --run_dir results_k4_unsup_ft --sweep snr --bnn_samples 10 --k_active 4Plot:
python scripts/plot_dnn_bnn.py --csv results_k4_unsup_ft/eval_active_users.csv --out_dir results_k4_unsup_ft/plots
python scripts/plot_dnn_bnn.py --csv results_k4_unsup_ft/eval_snr.csv --out_dir results_k4_unsup_ft/plots_snrTrain each method separately:
python scripts/train_unfolded.py --method deep_unfolded --out_dir results_unfolded --k_train 4 --steps 8000
python scripts/train_unfolded.py --method nesterov_unfolded --out_dir results_unfolded --k_train 4 --steps 8000
python scripts/train_unfolded.py --method bayes_nesterov_unfolded --out_dir results_unfolded --k_train 4 --steps 8000 --kl_beta 1e-7Compares truncated WMMSE (default 3 iterations) against all three trained unfolded models:
python scripts/evaluate_unfolded.py --run_dir results_unfolded --k_test_min 1 --k_test_max 8
python scripts/plot_unfolded.py --csv results_unfolded/eval_unfolded.csv --out_dir results_unfolded/plots_unfolded- Classical WMMSE (
bnn_wmmse/algorithms/wmmse.py): full iterative solver with power bisection; used as dataset labels and evaluation baseline. - Truncated WMMSE: same solver with fewer iterations (
--truncated_iters); baseline for unfolding experiments. - Deep unfolding (
DeepUnfoldedWMMSE): unrolls PGD steps with learnable step sizes. - Nesterov unfolding (
NesterovUnfoldedWMMSE): adds look-ahead and momentum on top of PGD unfolding. - Bayesian Nesterov unfolding (
BayesianNesterovUnfoldedWMMSE): step sizes and momenta are variational; MC samples at train/eval time. - DNN / BNN (
models/dnn_bnn.py): feedforward MLP maps(H, mask) → V; BNN uses Bayesian weights (separate from Bayesian unfolding).
- Run all commands from the repo root.
- Unsupervised training generates fresh random channels each step; supervised training reads from a saved
.npzdataset. - Experiment outputs go under
results_*/orruns_*/(gitignored).