Functional Weave
Code in Rust

math.fractional-power@1.0.0

impl/python.py

4,248 bytes · the Python implementation · view raw

import re
from typing import Optional

#: Every value is held as an integer count of 10^-18.
FIXED_SCALE = 10 ** 18

#: Intermediate values beyond 2^512 (about 10^136 after scaling) are refused.
_MAX_BITS = 512

_DECIMAL = re.compile(r"^[0-9]+(\.[0-9]{1,18})?$")


def _too_large(value: int) -> bool:
    return value.bit_length() > _MAX_BITS


def mul_fixed(a: int, b: int) -> int:
    """Product of two non-negative fixed-point values, floored."""
    return (a * b) // FIXED_SCALE


def pow_fixed_bounded(x: int, n: int, bound: Optional[int]) -> int:
    """x^n for a non-negative fixed-point x, by square-and-multiply from the
    lowest bit, flooring after every product.

    The order of operations is part of the contract: it is what makes
    TypeScript, Python and Rust agree to the last digit. Returns bound + 1 as
    soon as the result must exceed ``bound`` (x >= 1 only grows), so a
    bisection never builds an astronomically large number.
    """
    result = FIXED_SCALE
    base = x
    e = n
    while e > 0:
        if e % 2 == 1:
            result = mul_fixed(result, base)
            if bound is not None and result > bound and x >= FIXED_SCALE:
                return bound + 1
        e //= 2
        if e > 0:
            base = mul_fixed(base, base)
            if bound is not None and base > bound and x >= FIXED_SCALE:
                return bound + 1
            if _too_large(base):
                raise ValueError("fractional power result too large")
    if _too_large(result):
        raise ValueError("fractional power result too large")
    return result


def pow_fixed(x: int, n: int) -> int:
    """x^n in fixed point; n is a whole number of 0 or more."""
    return pow_fixed_bounded(x, n, None)


def root_fixed(x: int, q: int) -> int:
    """The q-th root of x in fixed point: the largest y with
    pow_fixed(y, q) <= x, found by bisection."""
    if q == 1:
        return x
    lo = 0
    # For x >= 1, (1 + (x - 1)/q)^q >= x (Bernoulli), so the root is below it.
    hi = FIXED_SCALE + (x - FIXED_SCALE) // q + 2 if x >= FIXED_SCALE else FIXED_SCALE + 1
    while hi - lo > 1:
        mid = (lo + hi) // 2
        if pow_fixed_bounded(mid, q, x) <= x:
            lo = mid
        else:
            hi = mid
    return lo


def fractional_power_fixed(x: int, p: int, q: int) -> int:
    """x^(p/q) in fixed point: the q-th root first, then the power, then the
    reciprocal if p < 0."""
    if isinstance(p, bool) or not isinstance(p, int) or p < -100000 or p > 100000:
        raise ValueError("exponentNumerator must be between -100000 and 100000, received %r" % (p,))
    if isinstance(q, bool) or not isinstance(q, int) or q < 1 or q > 100000:
        raise ValueError("exponentDenominator must be between 1 and 100000, received %r" % (q,))
    if x <= 0:
        raise ValueError("base must be greater than zero")
    root = root_fixed(x, q)
    if p >= 0:
        return pow_fixed(root, p)
    denominator = pow_fixed(root, -p)
    if denominator == 0:
        raise ValueError("fractional power result too large")
    result = (FIXED_SCALE * FIXED_SCALE) // denominator
    if _too_large(result):
        raise ValueError("fractional power result too large")
    return result


def parse_fixed(text: str) -> int:
    """Parse a non-negative decimal of at most 18 places into fixed point."""
    if not isinstance(text, str) or not _DECIMAL.fullmatch(text):
        raise ValueError('base must be a positive decimal with at most 18 places, received "%s"' % (text,))
    whole, _, fraction = text.partition(".")
    return int(whole) * FIXED_SCALE + int(fraction.ljust(18, "0"))


def fractional_power(base: str, exponent_numerator: int, exponent_denominator: int) -> str:
    """base^(exponent_numerator / exponent_denominator), to 12 decimal places.

    Computed in 18-place fixed point with a floor after every step, then
    rounded half-up to 12 places, so the last digit printed is right unless the
    true value sits within about 10^-15 of a rounding boundary.
    """
    x = parse_fixed(base)
    value = fractional_power_fixed(x, exponent_numerator, exponent_denominator)
    rounded = (value + 500000) // 1000000
    return "%d.%012d" % (rounded // 10 ** 12, rounded % 10 ** 12)