Xarray

Analyze labeled multi-dimensional data with xarray in your browser: select by label, resample time series, take weighted means and save netCDF files.

Python
# Working with multi-dimensional labeled data using xarray

import numpy as np
import pandas as pd
import xarray as xr
import matplotlib.pyplot as plt

# Generate sample data
times = pd.date_range('2023-01-01', periods=7)
latitudes = np.linspace(-90, 90, 10)
longitudes = np.linspace(-180, 180, 20)
temperature_data = 15 + 8 * np.random.randn(len(times), len(latitudes), len(longitudes))

# Create an xarray.Dataset
temperature_ds = xr.Dataset(
    {
        "temperature": (["time", "latitude", "longitude"], temperature_data)
    },
    coords={
        "time": times,
        "latitude": latitudes,
        "longitude": longitudes
    }
)

# Print the dataset
print(temperature_ds)

# Calculate the mean temperature over time
mean_temperature = temperature_ds.temperature.mean(dim="time")

# Plot the mean temperature
mean_temperature.plot(cmap="coolwarm")
plt.title("Mean Temperature Distribution")
plt.show()
<xarray.Dataset> Size: 11kB
Dimensions:      (time: 7, latitude: 10, longitude: 20)
Coordinates:
  * time         (time) datetime64[us] 56B 2023-01-01 2023-01-02 ... 2023-01-07
  * latitude     (latitude) float64 80B -90.0 -70.0 -50.0 ... 50.0 70.0 90.0
  * longitude    (longitude) float64 160B -180.0 -161.1 -142.1 ... 161.1 180.0
Data variables:
    temperature  (time, latitude, longitude) float64 11kB 22.9 18.64 ... 19.19

xarray puts names and labels on the dimensions of NumPy arrays. Instead of remembering that axis 0 is time, you write mean(dim="time") or sel(time="2023-01-03"). A DataArray holds one variable with its dimensions and coordinates, and a Dataset groups several variables that share them. Climate scientists, oceanographers and anyone with gridded data use it, and it reads and writes the netCDF files common in those fields. On this page it runs in your browser, so you can try it without installing anything. Run the example first, then paste any snippet below into a new cell to try it.

What the example does

The example makes a week of random daily temperatures on a grid of 10 latitudes by 20 longitudes. xr.Dataset() takes each variable as a pair of dimension names and values, here one variable, temperature, with the dimensions time, latitude and longitude. coords attaches the dates and grid positions as labels. Printing the Dataset lists the dimension sizes, the coordinates and the data variables. temperature_ds.temperature returns the variable as a DataArray, and .mean(dim="time") averages the seven days into a 10 × 20 grid. .plot() draws a 2D array as a color mesh with a color bar and labels the axes from the coordinate names. The data comes from np.random.randn() without a seed, so the map changes on every run.

Select data by label or position

sel() picks values by coordinate label and isel() by integer position. Both take dimension names, so the axis order doesn't matter:

import numpy as np
import pandas as pd
import xarray as xr

temps = xr.DataArray(
    np.arange(24.0).reshape(4, 3, 2),
    dims=("time", "lat", "lon"),
    coords={"time": pd.date_range("2024-01-01", periods=4),
            "lat": [10, 20, 30], "lon": [100, 110]},
    name="temperature",
)

print(temps.sel(time="2024-01-02", lat=20).values)    # by label
print(temps.isel(time=0, lat=-1).values)              # by position
print(temps.sel(lat=24, method="nearest").lat.item()) # closest label
print(temps.sel(time=slice("2024-01-02", "2024-01-03")).sizes)

Resample and group a time series

resample() changes the time step, like its pandas namesake, and groupby("time.season") groups dates by season. The data is a year of daily temperatures with a seasonal cycle and seeded noise:

import numpy as np
import pandas as pd
import xarray as xr
import matplotlib.pyplot as plt

rng = np.random.default_rng(seed=0)
days = pd.date_range("2024-01-01", "2024-12-31", freq="D")
seasonal = 12 - 10 * np.cos(2 * np.pi * days.dayofyear / 366)
temp = xr.DataArray(seasonal + rng.normal(0, 2, days.size),
                    coords={"time": days}, dims="time", name="temp")

monthly = temp.resample(time="MS").mean()
print(monthly.round(1).values)
print(temp.groupby("time.season").mean().round(1))

monthly.plot(marker="o")
plt.show()

"MS" means month start. The seasons come out in alphabetical order: DJF, JJA, MAM, SON.

Take an area-weighted mean

On a latitude-longitude grid, cells near the poles cover less area, so a plain mean gives them too much weight. weighted() fixes that. The last lines subtract each latitude's mean, which xarray lines up by dimension name:

import numpy as np
import xarray as xr

lat = np.arange(-80, 81, 20)
lon = np.arange(0, 360, 45)
rng = np.random.default_rng(seed=1)
base = 28 * np.cos(np.deg2rad(lat))[:, None] - 5          # warm equator
temp = xr.DataArray(base + rng.normal(0, 1, (lat.size, lon.size)),
                    coords={"lat": lat, "lon": lon}, dims=("lat", "lon"))

weights = np.cos(np.deg2rad(temp.lat))   # cells shrink toward the poles
print("plain mean:   ", round(float(temp.mean()), 2))
print("weighted mean:", round(float(temp.weighted(weights).mean()), 2))

anomaly = temp - temp.mean("lon")        # matched by dimension name
print(anomaly.dims, anomaly.sel(lat=0).round(2).values)

Save and open a netCDF file

xarray writes netCDF files with scipy or netCDF4, and neither is installed with it here. Importing scipy first avoids ValueError: cannot read or write netCDF files without netCDF4-python or scipy installed:

import scipy  # xarray writes .nc files with scipy or netCDF4
import xarray as xr

ds = xr.Dataset(
    {"temp": (("station", "month"), [[-4.3, 17.0], [8.0, 25.4]]),
     "rain": (("station", "month"), [[49, 81], [67, 19]])},
    coords={"station": ["north", "south"], "month": [1, 7]},
    attrs={"note": "made-up readings"},
)
ds.to_netcdf("readings.nc")

loaded = xr.open_dataset("readings.nc")
print(loaded)
print(loaded.to_dataframe())

open_dataset() loads values only when you use them, so the printout shows ... in place of the data. to_dataframe() turns the Dataset into a pandas DataFrame with one row per station and month.

Good to know

  • sel() needs an exact label. In the first snippet, temps.sel(lat=24) raises a KeyError. Pass method="nearest" to get the closest value.
  • Arithmetic between two arrays first aligns their labels and keeps only the labels both share. Adding arrays with x = 0, 1, 2 and x = 1, 2, 3 gives a result with only x = 1 and 2, and no error.