Note
Go to the end to download the full example code.
Riemannian GD vs. Adam#
Compare RiemannianGD vs RiemannianAdam on QFTBasis.
Run: python examples/optimizer_benchmark.py

Training with RiemannianGD (lr=0.01)...
0.66s, final loss 9.7034
Training with RiemannianAdam (lr=0.01)...
1.25s, final loss 8.9574
wrote: out/optimizer_benchmark.png
from __future__ import annotations
import time
from pathlib import Path
import jax
import jax.numpy as jnp
import pdft
from pdft.viz import TrainingHistory, plot_training_comparison
def main(out_dir: str | Path = "out") -> None:
out = Path(out_dir)
out.mkdir(parents=True, exist_ok=True)
target = jax.random.normal(jax.random.PRNGKey(3), (4, 4)).astype(jnp.complex128)
histories: list[TrainingHistory] = []
for label, opt in [
("RiemannianGD (lr=0.01)", pdft.RiemannianGD(lr=0.01)),
("RiemannianAdam (lr=0.01)", pdft.RiemannianAdam(lr=0.01)),
]:
basis = pdft.QFTBasis(m=2, n=2)
print(f"Training with {label}...")
t0 = time.perf_counter()
result = pdft.train_basis(
basis,
target=target,
loss=pdft.L1Norm(),
optimizer=opt,
steps=50,
seed=0,
)
elapsed = time.perf_counter() - t0
print(f" {elapsed:.2f}s, final loss {result.loss_history[-1]:.4f}")
histories.append(TrainingHistory(losses=result.loss_history, label=label))
fig_path = out / "optimizer_benchmark.png"
plot_training_comparison(histories, output_path=fig_path,
title="RiemannianGD vs RiemannianAdam on QFTBasis (L1 loss)")
print(f"wrote: {fig_path}")
if __name__ == "__main__":
main()
Total running time of the script: (0 minutes 2.110 seconds)