Waveguide with Dispersion
In this chapter, we'll extend our waveguide model to include dispersion effects. We'll see how to model wavelength-dependent behavior in optical waveguides.
Waveguide Model with 1st-order Dispersion
Parameters:
- \(\lambda\): Wavelength (µm)
- \(\lambda_0\): Reference wavelength (µm)
- \(n_{\text{eff}}\): Effective index at \(\lambda_0\)
- \(n_g\): Group index
- \(L\): Waveguide length (µm)
- \(\alpha\): Propagation loss (dB/cm)
Effective index with dispersion:
Transmission:
```python import jax.numpy as jnp from rich import print as rprint import sax
def waveguide_1st_order_dispersion( wl: float = 1.55, wl0: float = 1.55, neff: float = 2.34, ng: float = 3.4, length: float = 10.0, loss_db_per_cm: float = 0.5, ) -> sax.SDict: """A simple straight waveguide model.
Args:
wl: wavelength in microns.
wl0: reference wavelength in microns.
neff: effective index.
ng: group index.
length: length of the waveguide in microns.
loss_db_per_cm: loss in dB/cm.
"""
dwl = wl - wl0
dneff_dwl = (ng - neff) / wl0
_neff = neff - dwl * dneff_dwl
loss_db_per_µm = loss_db_per_cm * 1e-4
phase = 2 * jnp.pi * _neff * length / wl
amplitude = jnp.asarray(10 ** (-loss_db_per_µm * length / 20), dtype=complex)
transmission = amplitude * jnp.exp(1j * phase)
return sax.reciprocal(
{
("o1", "o2"): transmission,
}
)
rprint(waveguide_1st_order_dispersion()) rprint(waveguide_1st_order_dispersion(wl=1.6)) ```
Note that we used sax.reciprocal to simplify the code, since we have S21 = S12.
Waveguide Model with Generalized Dispersion
```python from functools import cache import jax import jax.numpy as jnp import xarray as xr import sax from jaxtyping import Array, ArrayLike from jax import grad
@cache def _xarr_neff_strip(): return ( xr.open_dataarray("data/neff_te0.nc") .load() .expand_dims({"neff_te0": ["neff_te0"]}, -1) )
def _interpolate_xarray( xarr: xr.DataArray, target_var: str, **kwargs: sax.FloatArrayLike ) -> sax.FloatArray:
dims = [str(d) for d in xarr.coords if d != target_var]
missing = [d for d in dims if d not in kwargs]
if missing:
raise ValueError(f"Missing required interpolation inputs: {missing}")
arrays = [jnp.asarray(kwargs[dim]) for dim in dims]
broadcasted = jnp.broadcast_arrays(*arrays)
shape = broadcasted[0].shape
interp_args = {dim: arr.ravel() for dim, arr in zip(dims, broadcasted, strict=True)}
result = sax.interpolate_xarray(xarr, **interp_args)[target_var]
return result.reshape(shape)
@jax.jit def waveguide_generalized_dispersion( wl: float, length: float = 10.0, loss_db_per_cm: float = 0.5, ) -> sax.SDict: """Waveguide with interpolated neff."""
with jax.ensure_compile_time_eval():
xarr = _xarr_neff_strip()
neff = _interpolate_xarray(
xarr=xarr,
target_var="neff_te0",
wavelength=wl,
)
loss_db_per_µm = loss_db_per_cm * 1e-4
phase = 2 * jnp.pi * neff * length / wl
amplitude = jnp.asarray(10 ** (-loss_db_per_µm * length / 20), dtype=complex)
transmission = amplitude * jnp.exp(1j * phase)
return sax.reciprocal({("o1", "o2"): transmission})
waveguide_generalized_dispersion(1.55) ```
```python import matplotlib.pyplot as plt
wls = jnp.linspace(1.5, 1.6, 1000) transmission = waveguide_generalized_dispersion(wl=wls, length=100.0, loss_db_per_cm=1) s21 = transmission[("o1", "o2")]
plt.figure(figsize=(9, 3)) plt.subplot(1, 2, 1) plt.plot(wls, jnp.abs(s21)**2) plt.xlabel("Wavelength (µm)") plt.ylabel("Amplitude") plt.subplot(1, 2, 2) plt.plot(wls, jnp.angle(s21)) plt.xlabel("Wavelength (µm)") plt.ylabel("Phase (rad)") plt.show() ```