def run_scale_curve(
input_parquet: pathlib.Path,
report_path: pathlib.Path,
*,
rungs: tuple[int, ...] | str = DEFAULT_RUNGS,
pca_components: int = 100,
umap_dim: int = 8,
umap_n_neighbors: int | None = None,
cluster_viz_method: str = "pca",
min_class_size: int = 20,
seed: int = 13,
pca_seed: int = 13,
umap_seed: int | None = None,
batch_size: int = 1000,
n_sample: int = 3,
prefix: str = "res",
model: str | None = None,
layer: str | None = None,
max_scale: int = 60,
min_scale: int | None = None,
clustering_grid_step: int = 5,
negative_label: str = NEGATIVE_LABEL,
screener_kind: str = "lda",
drop_rare_entities: bool = False,
min_mentions_per_entity: int = 20,
max_mentions_per_entity: int | None = None,
max_mentions_negative: int | None = None,
mention_cap_seed: int = 13,
) -> ScaleCurve | None:
"""Sweep ``clustering_sample_rows`` and fit ``log(MCS*) ~ a + b·log(N)``.
Writes ``scale_curve.json`` (consumable by ``pelinker-fit scale_curve_path=…``),
a log-log figure, and per-sample grid rows under ``report_path``.
Returns the fitted curve, or ``None`` when too few rungs survived to fit one.
"""
console = Console(force_terminal=True, width=120, legacy_windows=False)
input_parquet = input_parquet.expanduser()
if not input_parquet.exists():
console.print(f"[red]Input parquet not found: {input_parquet}[/red]")
return None
try:
rung_sizes = parse_rungs(rungs)
except ValueError as exc:
console.print(f"[red]{exc}[/red]")
return None
resolved_model, resolved_layer = model or "", layer or ""
if not resolved_model or not resolved_layer:
parsed = parse_model_filename(input_parquet.name, prefix)
if parsed is None:
console.print(
f"[red]Cannot parse model/layer from {input_parquet.name!r}; "
f"pass --model and --layer[/red]"
)
return None
parsed_model, parsed_layer = parsed
# parse_model_filename yields an int for a numeric layer; layers are strings
# everywhere else (they may be comma-specs like "1,2").
resolved_model = resolved_model or str(parsed_model)
resolved_layer = resolved_layer or str(parsed_layer)
report_path = report_path.expanduser()
report_path.mkdir(parents=True, exist_ok=True)
detail_path = report_path / CLUSTERING_SEARCH_GRID_PER_SAMPLE_CSV_BASENAME
base_config = clustering_optimization_config_for_run(
min_class_size=min_class_size,
max_scale=max_scale,
min_scale=min_scale,
clustering_grid_step=clustering_grid_step,
seed=seed,
clustering_sample_rows=None,
batch_size=batch_size,
negative_label=negative_label,
screener_kind=screener_kind,
drop_rare_entities=drop_rare_entities,
min_mentions_per_entity=min_mentions_per_entity,
max_mentions_per_entity=max_mentions_per_entity,
max_mentions_negative=max_mentions_negative,
mention_cap_seed=mention_cap_seed,
)
viz_method = cluster_viz_method.lower()
if viz_method not in ("pca", "umap"):
console.print(
f"[red]cluster_viz_method must be 'pca' or 'umap', got {cluster_viz_method!r}[/red]"
)
return None
transform_config = TransformConfig(
pca_components=pca_components,
umap_components=umap_dim,
umap_n_neighbors=umap_n_neighbors,
cluster_viz_components=min(3, umap_dim),
cluster_viz_method=cast(Literal["pca", "umap"], viz_method),
pca_seed=pca_seed,
umap_seed=umap_seed,
)
console.print(f"[cyan]Loading mention frame from {input_parquet}[/cyan]")
base_frame = load_selection_frame(
file_path=input_parquet,
config=base_config,
show_embedding_read_progress=True,
)
if base_frame is None or len(base_frame) == 0:
console.print("[red]No mention rows after load filters[/red]")
return None
n_available = len(base_frame)
console.print(f"[green]Loaded {n_available:,} mention rows[/green]")
# A rung above the frame size would silently collapse onto the full frame and
# duplicate an existing rung, which fit_scale_curve then discards as leverage-free.
usable = tuple(r for r in rung_sizes if r < n_available)
if len(usable) < len(rung_sizes):
dropped = tuple(r for r in rung_sizes if r >= n_available)
console.print(
f"[yellow]Dropping {len(dropped)} rung(s) at or above the frame size "
f"({n_available:,}): {dropped}[/yellow]"
)
if len(usable) < MIN_RUNGS_FOR_FIT:
console.print(
f"[red]Only {len(usable)} usable rung(s) below {n_available:,} rows; "
f"need {MIN_RUNGS_FOR_FIT}. Use a larger parquet or smaller rungs.[/red]"
)
return None
collected: list[ScaleRung] = []
for rung_size in usable:
console.print(f"[cyan]Rung {rung_size:,} rows — {n_sample} draw(s)[/cyan]")
rung, reports = _evaluate_rung(
base_frame,
sample_rows=rung_size,
base_config=base_config,
transform_config=transform_config,
n_sample=n_sample,
selected_labels=None,
console=console,
)
if rung is None:
continue
collected.append(rung)
merge_new_frames_into_per_sample_grid_csv(
detail_path,
[
grid_export_rows_from_report(
report,
model=resolved_model,
# Keep rungs distinguishable inside the shared grid CSV schema.
layer=f"{resolved_layer}:rows{rung.n_rows_realized}",
sample_idx=sample_idx,
chosen_min_cluster_size=rung.chosen_min_cluster_size,
)
for sample_idx, report in reports
],
)
console.print(
f"[green]Rung {rung.n_rows_realized:,}: "
f"chosen min_cluster_size = {rung.chosen_min_cluster_size}[/green]"
)
if len(collected) < MIN_RUNGS_FOR_FIT:
console.print(
f"[red]Only {len(collected)} rung(s) completed; need "
f"{MIN_RUNGS_FOR_FIT} to fit a curve.[/red]"
)
return None
curve = fit_scale_curve(collected)
_render_rung_table(console, curve)
payload = {
"schema": SCALE_CURVE_SCHEMA,
"model": resolved_model,
"layer": resolved_layer,
"input_parquet": str(input_parquet.resolve()),
"n_rows_available": int(n_available),
"pca_components": int(pca_components),
"umap_dim": int(umap_dim),
"umap_n_neighbors": (
None if umap_n_neighbors is None else int(umap_n_neighbors)
),
"n_sample": int(n_sample),
"curve": curve.to_jsonable(),
}
json_path = report_path / SCALE_CURVE_JSON_BASENAME
json_path.write_text(json.dumps(payload, indent=2) + "\n", encoding="utf-8")
console.print(f"[green]✓[/green] Scale curve written to [cyan]{json_path}[/cyan]")
console.print(
f"[bold]log(MCS) = {curve.log_intercept:.4f} + "
f"{curve.log_slope:.4f}·log(N)[/bold] (R² = {curve.r_squared:.4f})"
)
_warn_on_weak_fit(console, curve)
# Imported here so the pure-math and orchestration paths stay importable without
# matplotlib; the CLI is the only caller that needs a figure.
from pelinker.plotting import plot_scale_curve
plot_scale_curve(curve, report_path / SCALE_CURVE_FIGURE_BASENAME)
return curve