# 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 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.
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.
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() 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.
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)
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.
sel() needs an exact label. In the first snippet, temps.sel(lat=24)
raises a KeyError. Pass method="nearest" to get the closest value.