Skip to content
Open
Show file tree
Hide file tree
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
47 changes: 47 additions & 0 deletions ogcore/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -1415,3 +1415,50 @@ def json_to_dict(json_text):
msg += bline + "\n"
raise ValueError(msg)
return ordered_dict


def D3_to_df(tpi_vars, var="c_path", start_year=2025):
r"""
Reshape a three dimensional (T x S x J) model output array into a
tidy Pandas DataFrame with a (year, age) MultiIndex and one column
per ability type j.

This is a convenience for using transition-path output in other
software (for example plotting or spreadsheets), where a panel
(long-on-index, wide-on-type) layout is easier to work with than a
raw 3D NumPy array.

Args:
tpi_vars (dict): dictionary of transition-path variables, as
returned by the model (for example the ``TPI_vars`` output)
var (str): key in ``tpi_vars`` of the array to reshape; the array
must be three dimensional with shape (T, S, J)
start_year (int): calendar year of period t=0, used to label the
year level of the index

Returns:
df (Pandas DataFrame): a DataFrame indexed by a (Year, Age)
MultiIndex with T*S rows and J columns, one per ability type

Raises:
ValueError: if ``var`` is not in ``tpi_vars`` or the selected
array is not three dimensional
"""
if var not in tpi_vars:
raise ValueError(f"'{var}' is not a key in tpi_vars.")
data = np.asarray(tpi_vars[var])
if data.ndim != 3:
raise ValueError(
f"'{var}' has {data.ndim} dimensions; D3_to_df expects a "
"three dimensional (T x S x J) array."
)
T, S, J = data.shape
idx_t = [f"{start_year + i}" for i in range(T)]
idx_s = [f"{i}" for i in range(S)]
idx_j = [f"{i}" for i in range(J)]
multi_idx = pd.MultiIndex.from_product(
[idx_t, idx_s], names=["Year", "Age"]
)
reshaped_data = data.reshape(T * S, J)
df = pd.DataFrame(reshaped_data, index=multi_idx, columns=idx_j)
return df
26 changes: 26 additions & 0 deletions tests/test_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -998,3 +998,29 @@ def test_params_to_json_save(tmpdir):
with open(os.path.join(tmpdir, "test.json"), "r") as f:
j_str = f.read()
assert isinstance(j_str, str)


def test_D3_to_df():
"""D3_to_df reshapes a (T, S, J) array into a tidy (Year, Age) x J
DataFrame with the right shape, index, and values."""
T, S, J = 4, 3, 2
arr = np.arange(T * S * J, dtype=float).reshape(T, S, J)
tpi_vars = {"c_path": arr}
df = utils.D3_to_df(tpi_vars, var="c_path", start_year=2025)
# shape: T*S rows, J columns
assert df.shape == (T * S, J)
assert list(df.index.names) == ["Year", "Age"]
# first year label starts at start_year
assert df.index[0] == ("2025", "0")
# last year label is start_year + T - 1
assert df.index[-1] == (str(2025 + T - 1), str(S - 1))
# values round-trip: row (t, s), column j equals arr[t, s, j]
assert df.loc[("2026", "1")].tolist() == arr[1, 1, :].tolist()


def test_D3_to_df_errors():
"""D3_to_df raises on a missing key or a non-3D array."""
with pytest.raises(ValueError):
utils.D3_to_df({"c_path": np.zeros((2, 3, 2))}, var="missing")
with pytest.raises(ValueError):
utils.D3_to_df({"c_path": np.zeros((2, 3))}, var="c_path")
Loading