diff --git a/BMM/BMM_Theory/bmm_theory/io_csv.py b/BMM/BMM_Theory/bmm_theory/io_csv.py new file mode 100644 index 0000000..68c36e9 --- /dev/null +++ b/BMM/BMM_Theory/bmm_theory/io_csv.py @@ -0,0 +1,101 @@ +"""csv 输入输出层. + +case 输入 csv 列 (表头, 大小写不敏感, 缺省列取默认值): + case_id, batch_a, batch_b, m, n, k, + dtype_a, dtype_b, dtype_c, trans_a, trans_b, has_bias, deterministic_level + +方案输入 csv (评估模式): ImplPlan 全部字段, 见 ImplPlan.csv_fields(). + +输出 csv: + - 方案推荐模式: case 列 + plan_* 列 + 时延评估列 + feasible/bound_type/advice + - 方案评估模式: 同上, plan 来自用户输入, feasible/violations 反映约束校验 +""" + +from __future__ import annotations + +import csv +from pathlib import Path + +from .models import BmmCase, ImplPlan, EvalResult + + +_TRUE = {"1", "true", "yes", "y", "是"} + + +def _to_bool(v, default=False): + if v is None or str(v).strip() == "": + return default + return str(v).strip().lower() in _TRUE + + +def _to_int(v, default=0): + if v is None or str(v).strip() == "": + return default + return int(float(v)) + + +def _to_str(v, default=""): + if v is None: + return default + v = str(v).strip() + return v if v else default + + +def load_cases(path: str | Path) -> list: + """读取 case csv -> list[BmmCase].""" + cases = [] + with open(path, newline="", encoding="utf-8-sig") as f: + for i, row in enumerate(csv.DictReader(f)): + r = {k.strip().lower(): v for k, v in row.items() if k} + try: + cases.append(BmmCase( + case_id=_to_str(r.get("case_id"), f"case_{i}"), + batch_a=_to_int(r.get("batch_a"), 1), + batch_b=_to_int(r.get("batch_b"), 1), + m=_to_int(r.get("m"), 1), + n=_to_int(r.get("n"), 1), + k=_to_int(r.get("k"), 1), + dtype_a=_to_str(r.get("dtype_a"), "bf16"), + dtype_b=_to_str(r.get("dtype_b"), "bf16"), + dtype_c=_to_str(r.get("dtype_c"), "bf16"), + trans_a=_to_bool(r.get("trans_a")), + trans_b=_to_bool(r.get("trans_b")), + has_bias=_to_bool(r.get("has_bias")), + deterministic_level=_to_int(r.get("deterministic_level"), 0), + )) + except ValueError as e: + raise ValueError(f"{path} 第{i+2}行解析失败: {e}") from e + return cases + + +def load_plans(path: str | Path) -> list: + """读取方案 csv -> list[ImplPlan] (按 case_id 与 case 关联).""" + plans = [] + with open(path, newline="", encoding="utf-8-sig") as f: + for row in csv.DictReader(f): + r = {k.strip(): v for k, v in row.items() if k} + plans.append(ImplPlan.from_row(r)) + return plans + + +def save_results(path: str | Path, results: list) -> None: + """list[EvalResult] -> csv.""" + if not results: + return + rows = [r.to_row() for r in results] + fields = list(rows[0].keys()) + with open(path, "w", newline="", encoding="utf-8-sig") as f: + w = csv.DictWriter(f, fieldnames=fields) + w.writeheader() + w.writerows(rows) + + +def save_plans(path: str | Path, plans: list) -> None: + """list[ImplPlan] -> csv (标准结构体列).""" + if not plans: + return + with open(path, "w", newline="", encoding="utf-8-sig") as f: + w = csv.DictWriter(f, fieldnames=ImplPlan.csv_fields()) + w.writeheader() + for p in plans: + w.writerow(p.to_row())