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
60 changes: 56 additions & 4 deletions numpy_financial/_financial.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)

Expand Down
21 changes: 21 additions & 0 deletions numpy_financial/tests/test_financial.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down