Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
31 changes: 14 additions & 17 deletions docs/getting_started/install.md
Original file line number Diff line number Diff line change
Expand Up @@ -135,23 +135,8 @@ pip install "abtem[mps]"

The extra is restricted to macOS on Apple silicon, since Metal exists nowhere else.

To enable the backend, set `enable_mps` **before** *ab*TEM is imported — either through the environment:

```{code-block}
ABTEM_ENABLE_MPS=true python your_script.py
```

or in `~/.config/abtem/abtem.yaml`:

```{code-block}
enable_mps: true
```

```{note}
`abtem.config.set(enable_mps=True)` does **not** work: the setting decides whether PyTorch is imported before
FFTW, and by the time the call runs that has already been settled. PyTorch and FFTW each bundle their own copy
of `libomp`, and a process that loaded FFTW's first crashes inside ordinary PyTorch operations.
```
Nothing needs enabling: *ab*TEM imports PyTorch the first time the `mps` device is used, so importing *ab*TEM
stays as fast as without it.

You can verify that Metal support is available using the code below:

Expand All @@ -171,6 +156,18 @@ potential = abtem.Potential(atoms, gpts=512, device="mps")
probe = abtem.Probe(energy=200e3, semiangle_cutoff=20, device="mps")
```

or make Metal the default device, with `abtem.config.set({"device": "mps"})` or in `~/.config/abtem/abtem.yaml`:

```{code-block}
device: mps
```

```{note}
Import *ab*TEM (or PyTorch) before importing `pyfftw` yourself. PyTorch and FFTW each bundle their own copy of
`libomp`, and PyTorch's has to be loaded first; *ab*TEM arranges that when it is imported. If `pyfftw` was
imported earlier, using the `mps` device raises an error explaining this, rather than letting PyTorch crash.
```

Two things are worth knowing before you benchmark:

- **Metal is single precision.** `precision` must be `float32`; pairing `device="mps"` with `float64` is
Expand Down
Loading