Files
tianxuan/scripts/column_survey.py

100 lines
3.0 KiB
Python

#!/usr/bin/env python3
"""Analyze CSV columns and output a concise summary.
Usage:
runtime/python/python.exe scripts/column_survey.py --csv "data/*.csv"
"""
from __future__ import annotations
import argparse
import glob
import sys
from collections import OrderedDict
import polars as pl
def _analyze_columns(df: pl.DataFrame) -> list[OrderedDict]:
"""Return per-column stats: name, dtype, unique_count, null_pct, samples."""
rows: list[OrderedDict] = []
for col in df.columns:
series = df[col]
total = len(series)
null_cnt = series.null_count()
null_pct = (null_cnt / total * 100) if total > 0 else 0.0
unique_cnt = series.n_unique()
non_null = series.drop_nulls()
samples = non_null[:5].to_list() if len(non_null) > 0 else []
rows.append(OrderedDict([
("name", col),
("dtype", series.dtype),
("unique", unique_cnt),
("null_pct", null_pct),
("samples", samples),
("total_non_null", len(non_null)),
]))
return rows
def _fmt_samples(samples: list, total_non_null: int) -> str:
items = ", ".join(repr(s) for s in samples)
if 0 < len(samples) < total_non_null:
items += ", ..."
return f"[{items}]"
def main() -> None:
parser = argparse.ArgumentParser(description="Survey CSV columns")
parser.add_argument("--csv", required=True, help="Glob pattern for CSV files")
args = parser.parse_args()
files = sorted(glob.glob(args.csv, recursive=True))
if not files:
print(f"No files match pattern: {args.csv}", file=sys.stderr)
sys.exit(1)
# Load each file independently, group by schema
schema_groups: OrderedDict[str, list[pl.DataFrame]] = OrderedDict()
for f in files:
try:
lf = pl.scan_csv(f).head(5000)
df = lf.collect()
if df.is_empty():
print(f"Skip empty: {f}", file=sys.stderr)
continue
# Use schema string as grouping key
key = str(sorted(df.schema.items()))
schema_groups.setdefault(key, []).append(df)
except Exception as e:
print(f"Skip {f}: {e}", file=sys.stderr)
if not schema_groups:
print("No readable CSV data", file=sys.stderr)
sys.exit(1)
# Report each schema group
for idx, (schema_key, dfs) in enumerate(schema_groups.items()):
merged = pl.concat(dfs)
stats = _analyze_columns(merged)
if len(schema_groups) > 1:
group_label = f" (group {idx + 1}, {len(dfs)} file(s))"
else:
group_label = f" ({sum(len(v) for v in schema_groups.values())} file(s))"
print(f"=== Column Survey{group_label} ===")
for s in stats:
print(
f"{s['name']}: {s['dtype']}, "
f"unique={s['unique']}, "
f"null={s['null_pct']:.1f}%, "
f"sample={_fmt_samples(s['samples'], s['total_non_null'])}"
)
if __name__ == "__main__":
main()