"""Static vNext cost audit; no program simulation or optimality claims."""

import json
from itertools import product
from pathlib import Path

ROOT = Path(__file__).resolve().parent


def area(n, c, p, q, r, t, sh, h, qh, cache=0, bt=0):
    banks = 32 if sh else 0
    sm = (
        0.45
        + 0.018 * 32
        + 0.06 * 4
        + 0.18
        + c * (0.20 + 0.008 * p * q)
        + 0.018 * r * (1, 1.6, 2.6)[t - 1]
        + 0.004 * sh
        + 0.015 * banks
        + 0.10 * 2
        + 0.006 * 2 * 8
        + 0.514
        + 0.15
        + (0.20 + 0.004 * bt + 0.08 * c if bt else 0)
    )
    ca = 0.40 + 4.8 * cache + 0.12 * 8 + 0.0015 * 8 * 32 if cache else 0
    return n * sm + ca + 0.65 * h + 0.0015 * h * qh + 1 + 0.035 * n + 0.8 * 2**1.3 + 2


rows = []
# Deliberately bounded diagnostic menu, not the whole hardware space.
for n, c, shape, r, t, sh, h, qh, cache in product(
    (4, 8, 16, 24),
    (1, 2, 4),
    ((8, 8), (8, 16), (16, 16)),
    (32, 64, 128),
    (1, 2, 3),
    (0, 128),
    (2, 4, 8),
    (64, 128, 256),
    (0, 2),
):
    p, q = shape
    a = area(n, c, p, q, r, t, sh, h, qh, cache)
    # Roof ceilings, not a prediction: ignores fill/drain, accumulator reads,
    # finite workload, bank conflicts, memory traffic and instruction issue.
    read_bw = (128, 256, 512)[t - 1]
    feed_util = min(1, read_bw / (4 * (p + q) * c))
    rows.append(
        dict(
            n=n,
            c=c,
            shape=shape,
            rf=r,
            rf_t=t,
            sh=sh,
            h=h,
            qh=qh,
            cache=cache,
            area=a,
            feed_util=feed_util,
            nominal_tflops=n * c * p * q / 1000,
            rf_feed_ceiling_tflops=n * c * p * q / 1000 * feed_util,
        )
    )

summary = {}
for budget in (80, 100, 120):
    feasible = [x for x in rows if x["area"] <= budget]
    summary[budget] = dict(
        count=len(feasible),
        max_nominal_tflops=max(x["nominal_tflops"] for x in feasible),
        max_rf_feed_ceiling_tflops=max(x["rf_feed_ceiling_tflops"] for x in feasible),
        fraction_rf_feed_limited=sum(x["feed_util"] < 1 for x in feasible) / len(feasible),
        max_cache_mib=max(x["cache"] for x in feasible),
    )

queues = [
    dict(latency=latency, q=q, ideal_fraction=min(1, q * 64 / (latency + 2) / 32))
    for latency, q in product((150, 250, 400), (64, 128, 256))
]
# Explicit same-path partial energy tally for fully coalesced streaming read -> RF ->
# vector FMA. Cache bypass; one 4B weight -> one FMA = 2 FLOPs.
# Per 64B: 8B request + 64B response; 16 address elements; one front record;
# 8pJ control; RF write/read once; 16 vector FMA. Other operands excluded.
rf_price = 0.42 * 1.1
extra_pj_b = 72 * 2 / 64 + 8 / 64 + 16 * 0.5 / 64 + 2 / 64 + 2 * rf_price + 16 * 4 / 64
power = []
for e, h in product((60, 100, 150), (4, 8)):
    # A=100 intentionally conservative idle power for an area-feasible chip.
    base = 0.025 * 100 + 0.15 * h
    bw = h * 16e9
    p_partial = base + bw * (e + extra_pj_b) * 1e-12
    power.append(
        dict(
            e_hbm=e,
            channels=h,
            base_w=base,
            full_bandwidth_partial_path_power_at_100au_w=p_partial,
            allowed_gbs_at_20w=(20 - base) / (e + extra_pj_b) * 1000,
            allowed_gbs_at_26w=(26 - base) / (e + extra_pj_b) * 1000,
        )
    )

payload = dict(
    scope="static diagnostic subset; no simulator, no global search",
    sampled_menu_count=len(rows),
    area_summary=summary,
    queue_roof=queues,
    streaming_extra_pj_b=extra_pj_b,
    partial_path_power=power,
    dominance_example=dict(
        vector64_area=0.018 * 64,
        tc8x8_area=0.2 + 0.008 * 64,
        cache2_area=0.4 + 4.8 * 2 + 0.12 * 8 + 0.0015 * 8 * 32,
        rf64_low_area=0.018 * 64,
        rf64_high_area=0.018 * 64 * 2.6,
        sh64_16bank_area=0.004 * 64 + 0.015 * 16,
    ),
)
(ROOT / "audit.json").write_text(json.dumps(payload, indent=2) + "\n")
print(json.dumps(payload, indent=2))
