#!/usr/bin/env python3
"""Calculate a weight-storage lower bound, never a GPU training fit prediction.

Standard library only. No model is loaded or trained. The base payload assumes
all base parameters use exactly four bits, without quantization metadata.
Ordinary, unshared LoRA adapters are counted from explicit linear-layer shapes;
biases, embeddings, DoRA, gradients, optimizer states and activations are absent.
"""

import argparse
from dataclasses import dataclass
from decimal import Decimal, InvalidOperation

GIB = 1024 ** 3


def require_positive_integer(name, value):
    if type(value) is not int or value <= 0:
        raise ValueError(f"{name} must be a positive integer")


def ideal_four_bit_bytes(parameters):
    require_positive_integer("parameters", parameters)
    # Two four-bit values per byte, rounded up for an odd parameter count.
    return (parameters + 1) // 2


@dataclass(frozen=True)
class Projection:
    in_features: int
    out_features: int
    count: int

    def __post_init__(self):
        for name in ("in_features", "out_features", "count"):
            require_positive_integer(name, getattr(self, name))


def lora_parameter_count(projection, rank):
    require_positive_integer("rank", rank)
    return projection.count * rank * (
        projection.in_features + projection.out_features
    )


@dataclass(frozen=True)
class WeightFloor:
    ideal_base_bytes: int
    adapter_parameters: int
    adapter_weight_bytes: int

    @property
    def total_bytes(self):
        return self.ideal_base_bytes + self.adapter_weight_bytes

    def budget_status(self, budget_bytes):
        require_positive_integer("budget_bytes", budget_bytes)
        if self.total_bytes > budget_bytes:
            return "exceeds budget before training overhead"
        return "undetermined: training overhead is not estimated"


def estimate(parameters, projections, rank, adapter_bytes=4):
    """Assume all supplied weights are resident together, without offloading.

    adapter_bytes is an explicit storage assumption, not inferred from a
    compute dtype. The default is four bytes; inspect a real model's dtypes.
    Each counted projection must be a distinct, unshared base weight matrix.
    """
    base_bytes = ideal_four_bit_bytes(parameters)
    require_positive_integer("rank", rank)
    if type(adapter_bytes) is not int or adapter_bytes not in (2, 4, 8):
        raise ValueError("adapter_bytes must be 2, 4 or 8")
    projections = tuple(projections)
    if not projections or not all(isinstance(p, Projection) for p in projections):
        raise ValueError("provide at least one Projection")
    target_weights = sum(p.count * p.in_features * p.out_features for p in projections)
    if target_weights > parameters:
        raise ValueError("target matrices contain more weights than the base parameter count")
    adapters = sum(lora_parameter_count(p, rank) for p in projections)
    return WeightFloor(base_bytes, adapters, adapters * adapter_bytes)


def positive_integer(text):
    try:
        value = int(text)
        require_positive_integer("value", value)
        return value
    except ValueError as error:
        raise argparse.ArgumentTypeError("expected a positive integer") from error


def projection_argument(text):
    try:
        parts = text.split(":")
        if len(parts) != 3:
            raise ValueError("expected three fields")
        return Projection(*(int(part) for part in parts))
    except ValueError as error:
        raise argparse.ArgumentTypeError(
            "projection must be IN:OUT:COUNT with three positive integers"
        ) from error


def gib_argument(text):
    try:
        value = Decimal(text)
        if not value.is_finite() or value <= 0:
            raise ValueError("budget must be finite and positive")
        # A fractional byte of capacity is unusable, so round down.
        size = int(value * GIB)
        require_positive_integer("budget_bytes", size)
        return size
    except (ValueError, InvalidOperation, OverflowError) as error:
        raise argparse.ArgumentTypeError("expected a positive, finite GiB budget") from error


def main(argv=None):
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--parameters", required=True, type=positive_integer)
    parser.add_argument("--rank", required=True, type=positive_integer)
    parser.add_argument("--projection", required=True, action="append", type=projection_argument,
                        metavar="IN:OUT:COUNT")
    parser.add_argument("--adapter-bytes", type=int, choices=(2, 4, 8), default=4)
    parser.add_argument("--budget-gib", type=gib_argument)
    args = parser.parse_args(argv)
    try:
        result = estimate(args.parameters, args.projection, args.rank, args.adapter_bytes)
    except ValueError as error:
        parser.error(str(error))
    print(f"Ideal 4-bit base payload: {result.ideal_base_bytes:,} bytes ({result.ideal_base_bytes / GIB:.3f} GiB)")
    print(f"LoRA parameters: {result.adapter_parameters:,}")
    print(f"Adapter weights ({args.adapter_bytes} bytes each): {result.adapter_weight_bytes:,} bytes ({result.adapter_weight_bytes / GIB:.3f} GiB)")
    print(f"Weight-storage floor: {result.total_bytes:,} bytes ({result.total_bytes / GIB:.3f} GiB)")
    print("Excluded: quantization metadata, unquantized-weight uplift, gradients,")
    print("optimizer states, activations, temporary buffers, allocator/runtime overhead.")
    if args.budget_gib is not None:
        print(f"Budget verdict: {result.budget_status(args.budget_gib)}")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
