#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
fft_benchmark.py

CPU (SciPy FFT) と CUDA (CuPy/cuFFT) の FFT 性能を比較する教育用ベンチマーク。

比較する時間
------------
1. cpu_1
   SciPy FFT, 1 worker

2. cpu_all
   SciPy FFT, 利用可能な CPU worker をすべて使用

3. cuda_fft
   入力データがすでに GPU メモリ上にある場合の FFT 時間
   （CPU<->GPU 転送時間を含まない）

4. cuda_total
   CPU -> GPU 転送 + FFT + GPU -> CPU 転送
   実際のアプリケーションに近い end-to-end 時間

特徴
----
- 1D / 2D / 3D FFT
- complex64 / complex128
- warm-up 後に複数回計測
- CUDA の非同期実行を synchronize() して正しく計時
- GPU メモリ不足を見積もって危険なケースを自動スキップ
- CSV 出力
- 実行時間と speedup のグラフを PNG 出力
- CPU/GPU の FFT 結果をサンプル点で検証

必要パッケージ
--------------
CPU:
    numpy
    scipy
    matplotlib

CUDA:
    cupy
    例: CUDA 12.x 環境では cupy-cuda12x

例
--
標準ベンチマーク:
    python fft_benchmark.py

3D FFT のみ:
    python fft_benchmark.py --dims 3

3D サイズを指定:
    python fft_benchmark.py --dims 3 --sizes3 64,128,192,256,320

complex64 のみ:
    python fft_benchmark.py --dtypes complex64

繰り返し回数を増やす:
    python fft_benchmark.py --repeat 10

グラフを作らない:
    python fft_benchmark.py --no-plot
"""

from __future__ import annotations

import argparse
import csv
import math
import os
import platform
import sys
import time
from dataclasses import dataclass
from pathlib import Path
from typing import Optional

import numpy as np
import scipy
from scipy import fft as sp_fft

try:
    import matplotlib.pyplot as plt
except Exception:
    plt = None

try:
    import cupy as cp
except Exception:
    cp = None


DEFAULT_SIZES = {
    1: [2**16, 2**18, 2**20, 2**22],
    2: [256, 512, 1024, 2048],
    3: [64, 128, 192, 256],
}


@dataclass
class Result:
    dimension: int
    n: int
    shape: str
    points: int
    dtype: str
    input_mib: float
    cpu_1_s: float
    cpu_1_std_s: float
    cpu_all_s: float
    cpu_all_std_s: float
    cuda_fft_s: float
    cuda_fft_std_s: float
    cuda_total_s: float
    cuda_total_std_s: float
    speedup_cuda_fft_vs_cpu_all: float
    speedup_cuda_total_vs_cpu_all: float
    max_relative_error_sample: float
    note: str


def parse_int_list(text: str) -> list[int]:
    return [int(x.strip()) for x in text.split(",") if x.strip()]


def parse_str_list(text: str) -> list[str]:
    return [x.strip() for x in text.split(",") if x.strip()]


def shape_for(dim: int, n: int) -> tuple[int, ...]:
    return (n,) * dim


def shape_text(shape: tuple[int, ...]) -> str:
    return "x".join(str(v) for v in shape)


def mean_std(values: list[float]) -> tuple[float, float]:
    if not values:
        return math.nan, math.nan
    a = np.asarray(values, dtype=np.float64)
    return float(a.mean()), float(a.std(ddof=1)) if len(a) > 1 else 0.0


def safe_ratio(a: float, b: float) -> float:
    if not np.isfinite(a) or not np.isfinite(b) or b <= 0:
        return math.nan
    return a / b


def make_complex_array(
    shape: tuple[int, ...], dtype: np.dtype, seed: int
) -> np.ndarray:
    """
    大きな一時 complex128 配列を作らないよう、complex 配列を先に確保して
    real/imag を別々に代入する。
    """
    rng = np.random.default_rng(seed)
    x = np.empty(shape, dtype=dtype)

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

    # 一時配列は実数 1 個分だけに抑える。
    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 benchmark_cpu(
    x: np.ndarray,
    workers: int,
    repeat: int,
    warmup: int,
) -> tuple[float, float, np.ndarray]:
    y = None

    for _ in range(warmup):
        y = sp_fft.fftn(x, workers=workers)

    times = []
    for _ in range(repeat):
        t0 = time.perf_counter()
        y = sp_fft.fftn(x, workers=workers)
        t1 = time.perf_counter()
        times.append(t1 - t0)

    mean_s, std_s = mean_std(times)
    assert y is not None
    return mean_s, std_s, y


def clear_cupy_caches() -> None:
    if cp is None:
        return

    try:
        cp.get_default_memory_pool().free_all_blocks()
    except Exception:
        pass

    try:
        cp.get_default_pinned_memory_pool().free_all_blocks()
    except Exception:
        pass

    try:
        cp.fft.config.get_plan_cache().clear()
    except Exception:
        pass


def gpu_memory_ok(
    shape: tuple[int, ...],
    dtype: np.dtype,
    memory_fraction: float,
    workspace_factor: float = 3.0,
) -> tuple[bool, str]:
    """
    FFT 用の厳密な workspace はサイズ・GPU・cuFFT 実装で異なるため、
    input bytes の workspace_factor 倍を概算必要量として判定する。
    """
    if cp is None:
        return False, "CuPy not available"

    points = int(np.prod(shape))
    input_bytes = points * dtype.itemsize
    estimated = int(input_bytes * workspace_factor)

    try:
        free_b, total_b = cp.cuda.runtime.memGetInfo()
    except Exception as exc:
        return True, f"GPU memory query failed: {exc}"

    allowed = int(free_b * memory_fraction)

    if estimated > allowed:
        msg = (
            f"estimated GPU memory {estimated / 2**20:.0f} MiB > "
            f"{memory_fraction:.0%} of free memory "
            f"({allowed / 2**20:.0f} MiB)"
        )
        return False, msg

    return True, ""


def benchmark_cuda(
    x: np.ndarray,
    repeat: int,
    warmup: int,
    validate_points: int,
    cpu_reference: np.ndarray,
) -> tuple[float, float, float, float, float]:
    """
    Returns
    -------
    cuda_fft_mean, cuda_fft_std,
    cuda_total_mean, cuda_total_std,
    max_relative_error_sample
    """
    if cp is None:
        return (math.nan,) * 5

    # ---------- FFT only ----------
    x_gpu = cp.asarray(x)

    y_gpu = None
    for _ in range(warmup):
        y_gpu = cp.fft.fftn(x_gpu)
        cp.cuda.Stream.null.synchronize()

    fft_times = []
    for _ in range(repeat):
        cp.cuda.Stream.null.synchronize()
        t0 = time.perf_counter()
        y_gpu = cp.fft.fftn(x_gpu)
        cp.cuda.Stream.null.synchronize()
        t1 = time.perf_counter()
        fft_times.append(t1 - t0)

    cuda_fft_mean, cuda_fft_std = mean_std(fft_times)

    # ---------- CPU -> GPU -> FFT -> CPU ----------
    total_times = []
    y_host = None

    for _ in range(warmup):
        xg = cp.asarray(x)
        yg = cp.fft.fftn(xg)
        y_host = cp.asnumpy(yg)
        cp.cuda.Stream.null.synchronize()
        del xg, yg

    for _ in range(repeat):
        cp.cuda.Stream.null.synchronize()
        t0 = time.perf_counter()
        xg = cp.asarray(x)
        yg = cp.fft.fftn(xg)
        y_host = cp.asnumpy(yg)
        cp.cuda.Stream.null.synchronize()
        t1 = time.perf_counter()
        total_times.append(t1 - t0)
        del xg, yg

    cuda_total_mean, cuda_total_std = mean_std(total_times)

    # ---------- correctness check ----------
    # 全配列コピー・比較ではなく、ランダムなサンプル点だけ比較する。
    # y_gpu は pure CUDA FFT の最後の結果。
    err = math.nan

    if y_gpu is not None and validate_points > 0:
        total_points = x.size
        ns = min(validate_points, total_points)

        rng = np.random.default_rng(123456)
        idx = rng.choice(total_points, size=ns, replace=False)

        gpu_sample = cp.asnumpy(y_gpu.ravel()[idx])
        cpu_sample = cpu_reference.ravel()[idx]

        denominator = np.maximum(np.abs(cpu_sample), np.finfo(np.float64).eps)
        relative = np.abs(gpu_sample - cpu_sample) / denominator
        err = float(np.max(relative))

    del x_gpu
    if y_gpu is not None:
        del y_gpu
    if y_host is not None:
        del y_host

    clear_cupy_caches()

    return (
        cuda_fft_mean,
        cuda_fft_std,
        cuda_total_mean,
        cuda_total_std,
        err,
    )


def detect_gpu() -> tuple[bool, str]:
    if cp is None:
        return False, "CuPy is not installed"

    try:
        count = cp.cuda.runtime.getDeviceCount()
        if count < 1:
            return False, "No CUDA device found"

        dev = cp.cuda.Device(0)
        props = cp.cuda.runtime.getDeviceProperties(dev.id)
        raw_name = props["name"]
        name = raw_name.decode() if isinstance(raw_name, bytes) else str(raw_name)

        # 軽い CUDA 呼び出しで実際に動くことも確認
        a = cp.asarray([1.0], dtype=cp.float32)
        _ = a * 2
        cp.cuda.Stream.null.synchronize()

        return True, name
    except Exception as exc:
        return False, f"CUDA initialization failed: {exc}"


def system_info(gpu_ok: bool, gpu_name: str) -> dict[str, str]:
    info = {
        "python": sys.version.split()[0],
        "numpy": np.__version__,
        "scipy": scipy.__version__,
        "platform": platform.platform(),
        "processor": platform.processor() or "unknown",
        "logical_cpu_count": str(os.cpu_count()),
        "cupy": "not installed",
        "gpu": gpu_name if gpu_ok else "not available",
    }

    if cp is not None:
        info["cupy"] = cp.__version__

        try:
            info["cuda_runtime"] = str(cp.cuda.runtime.runtimeGetVersion())
        except Exception:
            info["cuda_runtime"] = "unknown"

        try:
            info["cuda_driver"] = str(cp.cuda.runtime.driverGetVersion())
        except Exception:
            info["cuda_driver"] = "unknown"

    return info


def print_system_info(info: dict[str, str]) -> None:
    print("=" * 72)
    print("FFT benchmark: CPU (SciPy) vs CUDA (CuPy/cuFFT)")
    print("=" * 72)

    for k, v in info.items():
        print(f"{k:20s}: {v}")

    print("=" * 72)


def write_csv(path: Path, results: list[Result], info: dict[str, str]) -> None:
    fieldnames = list(Result.__dataclass_fields__.keys())

    with path.open("w", newline="", encoding="utf-8-sig") as f:
        # コメント行として環境情報を保存
        for k, v in info.items():
            f.write(f"# {k}: {v}\n")

        writer = csv.DictWriter(f, fieldnames=fieldnames)
        writer.writeheader()

        for r in results:
            writer.writerow(r.__dict__)


def plot_results(results: list[Result], outdir: Path) -> None:
    if plt is None:
        print("matplotlib is not available: plots are skipped.")
        return

    dims = sorted(set(r.dimension for r in results))
    dtypes = sorted(set(r.dtype for r in results))

    for dim in dims:
        for dtype in dtypes:
            rows = [
                r for r in results
                if r.dimension == dim and r.dtype == dtype
            ]
            rows.sort(key=lambda r: r.points)

            if not rows:
                continue

            x = np.asarray([r.points for r in rows], dtype=np.float64)

            # ----- execution time -----
            fig, ax = plt.subplots(figsize=(8.0, 5.4))

            ax.plot(
                x,
                [r.cpu_1_s for r in rows],
                marker="o",
                label="CPU: SciPy, 1 worker",
            )
            ax.plot(
                x,
                [r.cpu_all_s for r in rows],
                marker="o",
                label="CPU: SciPy, all workers",
            )

            if any(np.isfinite(r.cuda_fft_s) for r in rows):
                ax.plot(
                    x,
                    [r.cuda_fft_s for r in rows],
                    marker="o",
                    label="CUDA: cuFFT only",
                )
                ax.plot(
                    x,
                    [r.cuda_total_s for r in rows],
                    marker="o",
                    label="CUDA: transfer + FFT + transfer",
                )

            ax.set_xscale("log")
            ax.set_yscale("log")
            ax.set_xlabel("Number of complex data points")
            ax.set_ylabel("Time [s]")
            ax.set_title(f"{dim}D FFT benchmark ({dtype})")
            ax.grid(True, which="both", alpha=0.3)
            ax.legend()
            fig.tight_layout()

            p = outdir / f"fft_time_{dim}d_{dtype}.png"
            fig.savefig(p, dpi=160)
            plt.close(fig)

            # ----- speedup -----
            if any(np.isfinite(r.cuda_fft_s) for r in rows):
                fig, ax = plt.subplots(figsize=(8.0, 5.4))

                ax.plot(
                    x,
                    [r.speedup_cuda_fft_vs_cpu_all for r in rows],
                    marker="o",
                    label="CPU(all) / CUDA FFT",
                )
                ax.plot(
                    x,
                    [r.speedup_cuda_total_vs_cpu_all for r in rows],
                    marker="o",
                    label="CPU(all) / CUDA total",
                )

                ax.axhline(1.0, linestyle="--", linewidth=1.0)
                ax.set_xscale("log")
                ax.set_yscale("log")
                ax.set_xlabel("Number of complex data points")
                ax.set_ylabel("Speedup")
                ax.set_title(
                    f"CUDA speedup relative to CPU(all) ({dim}D, {dtype})"
                )
                ax.grid(True, which="both", alpha=0.3)
                ax.legend()
                fig.tight_layout()

                p = outdir / f"fft_speedup_{dim}d_{dtype}.png"
                fig.savefig(p, dpi=160)
                plt.close(fig)


def format_time(x: float) -> str:
    if not np.isfinite(x):
        return "   N/A   "
    if x < 1e-3:
        return f"{x * 1e6:7.1f} us"
    if x < 1.0:
        return f"{x * 1e3:7.2f} ms"
    return f"{x:7.3f} s"


def run_case(
    dim: int,
    n: int,
    dtype_name: str,
    repeat: int,
    warmup: int,
    gpu_ok: bool,
    gpu_memory_fraction: float,
    validate_points: int,
    seed: int,
) -> Result:
    dtype = np.dtype(dtype_name)
    shape = shape_for(dim, n)
    points = int(np.prod(shape))
    input_mib = points * dtype.itemsize / 2**20

    print()
    print(
        f"[{dim}D] shape={shape_text(shape):>14s}  "
        f"dtype={dtype_name:10s}  "
        f"points={points:>12,d}  input={input_mib:8.1f} MiB"
    )

    note_parts: list[str] = []

    try:
        x = make_complex_array(shape, dtype, seed)
    except MemoryError:
        msg = "Host memory allocation failed"
        print("  SKIP:", msg)
        return Result(
            dim, n, shape_text(shape), points, dtype_name, input_mib,
            math.nan, math.nan, math.nan, math.nan,
            math.nan, math.nan, math.nan, math.nan,
            math.nan, msg,
        )

    # CPU 1 worker
    cpu_1_s, cpu_1_std, _ = benchmark_cpu(
        x, workers=1, repeat=repeat, warmup=warmup
    )

    # CPU all workers
    # scipy.fft の workers=-1 は利用可能な CPU worker をすべて使用する指定。
    cpu_all_s, cpu_all_std, cpu_ref = benchmark_cpu(
        x, workers=-1, repeat=repeat, warmup=warmup
    )

    cuda_fft_s = math.nan
    cuda_fft_std = math.nan
    cuda_total_s = math.nan
    cuda_total_std = math.nan
    error = math.nan

    if gpu_ok:
        ok, reason = gpu_memory_ok(
            shape,
            dtype,
            memory_fraction=gpu_memory_fraction,
        )

        if ok:
            try:
                (
                    cuda_fft_s,
                    cuda_fft_std,
                    cuda_total_s,
                    cuda_total_std,
                    error,
                ) = benchmark_cuda(
                    x=x,
                    repeat=repeat,
                    warmup=warmup,
                    validate_points=validate_points,
                    cpu_reference=cpu_ref,
                )
            except Exception as exc:
                note_parts.append(f"CUDA failed: {exc}")
                clear_cupy_caches()
        else:
            note_parts.append(f"CUDA skipped: {reason}")
    else:
        note_parts.append("CUDA unavailable")

    speed_fft = safe_ratio(cpu_all_s, cuda_fft_s)
    speed_total = safe_ratio(cpu_all_s, cuda_total_s)

    print(f"  CPU 1 worker : {format_time(cpu_1_s)}")
    print(f"  CPU all      : {format_time(cpu_all_s)}")

    if np.isfinite(cuda_fft_s):
        print(
            f"  CUDA FFT     : {format_time(cuda_fft_s)}"
            f"   speedup={speed_fft:7.2f} x"
        )
        print(
            f"  CUDA total   : {format_time(cuda_total_s)}"
            f"   speedup={speed_total:7.2f} x"
        )
        print(f"  check error  : {error:.3e}")
    else:
        print("  CUDA         : skipped / unavailable")

    return Result(
        dimension=dim,
        n=n,
        shape=shape_text(shape),
        points=points,
        dtype=dtype_name,
        input_mib=input_mib,
        cpu_1_s=cpu_1_s,
        cpu_1_std_s=cpu_1_std,
        cpu_all_s=cpu_all_s,
        cpu_all_std_s=cpu_all_std,
        cuda_fft_s=cuda_fft_s,
        cuda_fft_std_s=cuda_fft_std,
        cuda_total_s=cuda_total_s,
        cuda_total_std_s=cuda_total_std,
        speedup_cuda_fft_vs_cpu_all=speed_fft,
        speedup_cuda_total_vs_cpu_all=speed_total,
        max_relative_error_sample=error,
        note="; ".join(note_parts),
    )


def build_parser() -> argparse.ArgumentParser:
    p = argparse.ArgumentParser(
        description="CPU (SciPy FFT) vs CUDA (CuPy/cuFFT) benchmark"
    )

    p.add_argument(
        "--dims",
        default="1,2,3",
        help="FFT dimensions, e.g. 1,2,3  [default: 1,2,3]",
    )
    p.add_argument(
        "--sizes1",
        default=",".join(str(x) for x in DEFAULT_SIZES[1]),
        help="1D FFT lengths",
    )
    p.add_argument(
        "--sizes2",
        default=",".join(str(x) for x in DEFAULT_SIZES[2]),
        help="2D FFT side lengths; shape is n x n",
    )
    p.add_argument(
        "--sizes3",
        default=",".join(str(x) for x in DEFAULT_SIZES[3]),
        help="3D FFT side lengths; shape is n x n x n",
    )
    p.add_argument(
        "--dtypes",
        default="complex64,complex128",
        help="comma separated: complex64,complex128",
    )
    p.add_argument(
        "--repeat",
        type=int,
        default=5,
        help="number of timed repetitions [default: 5]",
    )
    p.add_argument(
        "--warmup",
        type=int,
        default=1,
        help="warm-up repetitions [default: 1]",
    )
    p.add_argument(
        "--validate-points",
        type=int,
        default=1024,
        help="number of FFT output points used for CPU/GPU validation",
    )
    p.add_argument(
        "--gpu-memory-fraction",
        type=float,
        default=0.75,
        help=(
            "fraction of currently free GPU memory allowed for the "
            "estimated FFT workload [default: 0.75]"
        ),
    )
    p.add_argument(
        "--seed",
        type=int,
        default=12345,
        help="random seed [default: 12345]",
    )
    p.add_argument(
        "--output",
        default="fft_benchmark.csv",
        help="CSV output path [default: fft_benchmark.csv]",
    )
    p.add_argument(
        "--plot-dir",
        default="fft_benchmark_plots",
        help="directory for PNG plots",
    )
    p.add_argument(
        "--no-plot",
        action="store_true",
        help="do not generate plots",
    )

    return p


def main() -> int:
    args = build_parser().parse_args()

    dims = parse_int_list(args.dims)
    dtypes = parse_str_list(args.dtypes)

    for dim in dims:
        if dim not in (1, 2, 3):
            raise SystemExit(f"Unsupported dimension: {dim}")

    for dtype in dtypes:
        if dtype not in ("complex64", "complex128"):
            raise SystemExit(f"Unsupported dtype: {dtype}")

    if args.repeat < 1:
        raise SystemExit("--repeat must be >= 1")

    if args.warmup < 0:
        raise SystemExit("--warmup must be >= 0")

    if not (0.05 <= args.gpu_memory_fraction <= 1.0):
        raise SystemExit("--gpu-memory-fraction must be in [0.05, 1.0]")

    sizes = {
        1: parse_int_list(args.sizes1),
        2: parse_int_list(args.sizes2),
        3: parse_int_list(args.sizes3),
    }

    gpu_ok, gpu_name = detect_gpu()
    info = system_info(gpu_ok, gpu_name)
    print_system_info(info)

    if not gpu_ok:
        print("CUDA benchmark will be skipped.")
        print(f"Reason: {gpu_name}")

    results: list[Result] = []

    case_id = 0
    for dim in dims:
        for dtype in dtypes:
            for n in sizes[dim]:
                case_id += 1

                result = run_case(
                    dim=dim,
                    n=n,
                    dtype_name=dtype,
                    repeat=args.repeat,
                    warmup=args.warmup,
                    gpu_ok=gpu_ok,
                    gpu_memory_fraction=args.gpu_memory_fraction,
                    validate_points=args.validate_points,
                    seed=args.seed + case_id,
                )
                results.append(result)

    output_path = Path(args.output)
    write_csv(output_path, results, info)

    print()
    print(f"CSV written: {output_path.resolve()}")

    if not args.no_plot:
        plot_dir = Path(args.plot_dir)
        plot_dir.mkdir(parents=True, exist_ok=True)
        plot_results(results, plot_dir)
        print(f"Plots written: {plot_dir.resolve()}")

    print()
    print("Interpretation:")
    print("  speedup > 1 : CUDA is faster than CPU(all workers)")
    print("  speedup < 1 : CPU(all workers) is faster")
    print()
    print(
        "cuda_fft excludes PCIe/system-memory transfer; "
        "cuda_total includes transfer."
    )
    print(
        "For small FFTs, transfer and launch overhead can dominate, "
        "so CPU may be faster."
    )

    return 0


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