diff --git a/ogcore/utils.py b/ogcore/utils.py index 073dc4dc6..a2055c172 100644 --- a/ogcore/utils.py +++ b/ogcore/utils.py @@ -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 diff --git a/tests/test_utils.py b/tests/test_utils.py index 2b301c72d..3a2ba788c 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -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")