#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
同一の 3D FFT を複数 backend で比較するサンプル。

Examples
--------
NumPy + CuPy:
    python examples/fft_compare.py

NumPy のみ:
    python examples/fft_compare.py --backends numpy

Intel/SYCL dpnp も比較:
    python examples/fft_compare.py --backends numpy,cupy,dpnp

dpnp の device selector:
    python examples/fft_compare.py --backends numpy,dpnp \
        --dpnp-device level_zero:gpu:0

大きさ:
    python examples/fft_compare.py --n 320
"""

from __future__ import annotations

import argparse
import time

import numpy as np

from tkcompute import get_backend


def make_input(n: int, dtype: np.dtype, seed: int = 12345) -> np.ndarray:
    """比較用の大きめ 3D complex ndarray を作る。"""
    shape = (n, n, n)
    rng = np.random.default_rng(seed)
    x = np.empty(shape, dtype=dtype)

    real_dtype = np.float32 if dtype == np.complex64 else np.float64

    tmp = rng.random(shape, dtype=real_dtype)
    x.real = tmp - real_dtype(0.5)
    del tmp

    tmp = rng.random(shape, dtype=real_dtype)
    x.imag = tmp - real_dtype(0.5)
    del tmp

    return x


def build_backend(name: str, args):
    if name == "numpy":
        return get_backend("numpy", workers=args.workers)

    if name == "cupy":
        return get_backend("cupy", device_id=args.cupy_device)

    if name == "dpnp":
        return get_backend("dpnp", device=args.dpnp_device)

    # registry に追加された custom backend も使えるようにする。
    return get_backend(name)


def benchmark_backend(bk, x0: np.ndarray):
    """
    transfer-in / FFT-only / transfer-out を分けて測る。

    NumPy backend の transfer は np.asarray() / np.asarray() なので、
    x0 が ndarray なら実質 no-copy に近い。
    """

    # ----- host -> backend -----
    t0 = time.perf_counter()
    x = bk.asarray(x0)
    bk.synchronize()
    transfer_in = time.perf_counter() - t0

    # warm-up
    _ = bk.fftn(x)
    bk.synchronize()

    # ----- FFT only -----
    t0 = time.perf_counter()
    y = bk.fftn(x)
    bk.synchronize()
    fft_time = time.perf_counter() - t0

    # ----- backend -> NumPy -----
    t0 = time.perf_counter()
    y_np = bk.to_numpy(y)
    bk.synchronize()
    transfer_out = time.perf_counter() - t0

    total = transfer_in + fft_time + transfer_out

    return {
        "result": y_np,
        "transfer_in_s": transfer_in,
        "fft_s": fft_time,
        "transfer_out_s": transfer_out,
        "total_s": total,
    }


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--n", type=int, default=256)
    p.add_argument(
        "--dtype",
        choices=("complex64", "complex128"),
        default="complex64",
    )
    p.add_argument(
        "--backends",
        default="numpy,cupy",
        help="comma-separated backend names",
    )
    p.add_argument(
        "--workers",
        type=int,
        default=-1,
        help="SciPy FFT workers for NumPy backend",
    )
    p.add_argument("--cupy-device", type=int, default=0)
    p.add_argument(
        "--dpnp-device",
        default="gpu",
        help='SYCL selector, e.g. "gpu", "cpu", "level_zero:gpu:0"',
    )
    args = p.parse_args()

    dtype = np.complex64 if args.dtype == "complex64" else np.complex128
    shape = (args.n,) * 3
    mib = np.prod(shape) * np.dtype(dtype).itemsize / 1024**2

    print("=" * 78)
    print("tkcompute backend comparison")
    print("=" * 78)
    print(f"shape       : {shape}")
    print(f"dtype       : {args.dtype}")
    print(f"input size  : {mib:.1f} MiB")
    print(f"backends    : {args.backends}")
    print("=" * 78)

    print("\nCreating input NumPy array ...")
    x0 = make_input(args.n, dtype)
    print("ready")

    requested = [
        x.strip().lower()
        for x in args.backends.split(",")
        if x.strip()
    ]

    results = {}
    reference = None

    for name in requested:
        print(f"\n[{name}]")

        try:
            bk = build_backend(name, args)
        except Exception as exc:
            print(f"SKIP: {exc}")
            continue

        print(f"backend      : {bk.name}")
        print(f"device       : {bk.get_device_name()}")

        info = bk.get_device_info()
        for key in sorted(info):
            if key not in ("backend", "device"):
                print(f"{key:12s} : {info[key]}")

        try:
            r = benchmark_backend(bk, x0)
        except Exception as exc:
            print(f"FAILED during calculation: {exc}")
            continue

        results[name] = r

        print(f"to backend   : {r['transfer_in_s']:.6f} s")
        print(f"FFT only     : {r['fft_s']:.6f} s")
        print(f"to NumPy     : {r['transfer_out_s']:.6f} s")
        print(f"total        : {r['total_s']:.6f} s")

        if reference is None:
            reference = r["result"]
            print("check        : reference")
        else:
            if args.dtype == "complex64":
                rtol, atol = 2e-4, 2e-3
            else:
                rtol, atol = 1e-10, 1e-8

            ok = np.allclose(reference, r["result"], rtol=rtol, atol=atol)
            max_abs = float(np.max(np.abs(reference - r["result"])))
            print(f"allclose     : {ok}")
            print(f"max abs err  : {max_abs:.6e}")

    if "numpy" in results:
        cpu_fft = results["numpy"]["fft_s"]
        cpu_total = results["numpy"]["total_s"]

        print("\n" + "=" * 78)
        print("Speedup relative to NumPy/SciPy")
        print("=" * 78)

        for name, r in results.items():
            if name == "numpy":
                continue
            print(
                f"{name:10s} "
                f"FFT-only={cpu_fft / r['fft_s']:8.2f} x   "
                f"total={cpu_total / r['total_s']:8.2f} x"
            )


if __name__ == "__main__":
    main()
