diff --git a/numpy_financial/_financial.py b/numpy_financial/_financial.py index 1adbbea..9df820a 100644 --- a/numpy_financial/_financial.py +++ b/numpy_financial/_financial.py @@ -13,7 +13,16 @@ from collections.abc import Iterable, Mapping, Sequence from decimal import Decimal -from typing import Any, Callable, Final, Literal, Protocol, TypeAlias, TypeVar, overload +from typing import ( + Any, + Callable, + Final, + Literal, + Protocol, + TypeAlias, + TypeVar, + overload, +) import numpy as np import numpy.typing as npt @@ -490,6 +499,46 @@ def _value_like(arr: npt.NDArray[Any], value: Decimal | float) -> Any: return Decimal(value) return np.array(value, dtype=arr.dtype).item(0) + +def _broadcast_payment_inputs( + rate: _ArrayLike, + per: _ArrayLike, + nper: _ArrayLike, + pv: _ArrayLike, + fv: _ArrayLike, + when: _ArrayLike, +): + """Broadcast row parameters over nested periods in object arrays.""" + period_values = np.asarray(per) + if period_values.ndim == 1 and period_values.dtype == object: + rows = [np.asarray(value) for value in period_values] + if ( + rows + and rows[0].ndim > 0 + and all(row.shape == rows[0].shape for row in rows) + ): + period_values = np.stack(rows) + row_count = period_values.shape[0] + + def expand_row_parameter(value: _ArrayLike) -> _ArrayLike: + if np.ndim(value) == 1 and np.shape(value) == (row_count,): + return np.asarray(value)[:, np.newaxis] + return value + + rate, nper, pv, fv, when = ( + expand_row_parameter(rate), + expand_row_parameter(nper), + expand_row_parameter(pv), + expand_row_parameter(fv), + expand_row_parameter(when), + ) + + rate, per, nper, pv, fv, when = np.broadcast_arrays( + rate, period_values, nper, pv, fv, when + ) + return rate, per, nper, pv, fv, when + + @overload def ipmt( rate: _AsFloat, @@ -611,9 +660,9 @@ def ipmt(rate, per, nper, pv, fv: Any = 0, when: _When = 'end') -> Any: np.float64(-112.98) """ - when = _convert_when(when) - rate, per, nper, pv, fv, when = np.broadcast_arrays(rate, per, nper, - pv, fv, when) + rate, per, nper, pv, fv, when = _broadcast_payment_inputs( + rate, per, nper, pv, fv, _convert_when(when) + ) total_pmt = pmt(rate, nper, pv, fv, when) ipmt_array = np.array(_rbl(rate, per, total_pmt, pv, when) * rate) @@ -741,6 +790,9 @@ def ppmt(rate, per, nper, pv, fv: Any = 0, when: _When = 'end'): pmt, pv, ipmt """ + rate, per, nper, pv, fv, when = _broadcast_payment_inputs( + rate, per, nper, pv, fv, _convert_when(when) + ) total = pmt(rate, nper, pv, fv, when) return total - ipmt(rate, per, nper, pv, fv, when) diff --git a/numpy_financial/tests/test_financial.py b/numpy_financial/tests/test_financial.py index 765e1b2..4136490 100644 --- a/numpy_financial/tests/test_financial.py +++ b/numpy_financial/tests/test_financial.py @@ -680,6 +680,27 @@ def test_0d_inputs(self): assert numpy.isscalar(npf.ipmt(*args)) + def test_nested_period_arrays_from_dataframe_column(self): + rate = numpy.array([0.05, 0.07]) + per = numpy.empty(2, dtype=object) + per[0] = numpy.arange(1, 5) + per[1] = numpy.arange(1, 5) + nper = numpy.array([10, 10]) + pv = numpy.array([10000, 12000]) + + per_matrix = numpy.stack(per) + expected_ipmt = npf.ipmt( + rate[:, numpy.newaxis], per_matrix, nper[:, numpy.newaxis], + pv[:, numpy.newaxis] + ) + expected_ppmt = npf.ppmt( + rate[:, numpy.newaxis], per_matrix, nper[:, numpy.newaxis], + pv[:, numpy.newaxis] + ) + assert_allclose(npf.ipmt(rate, per, nper, pv), expected_ipmt) + assert_allclose(npf.ppmt(rate, per, nper, pv), expected_ppmt) + + class TestFv: def test_float(self): assert_allclose(