HCAI-Lab/w2-consensus-deepdive-unlearning-artifacts / social-data-attribution-w2 /src /dolma /bin_analysis /stats.py
| """Statistics and mismatch detection for the 24x24 bin grid.""" | |
| from __future__ import annotations | |
| from pathlib import Path | |
| import polars as pl | |
| from dolma.constants import ( | |
| BIN_EMPTY, | |
| BIN_FULL, | |
| BIN_PARTIAL, | |
| BIN_SPARSE, | |
| TARGET_DOCS_PER_BIN, | |
| TOKEN_FLOOR_PER_BIN, | |
| ) | |
| from dolma.manifest_paths import scan_manifest_polars | |
| DEFAULT_DOC_THRESHOLD = TARGET_DOCS_PER_BIN | |
| DEFAULT_TOKEN_FLOOR = TOKEN_FLOOR_PER_BIN | |
| DEFAULT_TOKEN_CAP = 5_000_000 | |
| def load_manifest(manifest_path: Path) -> pl.LazyFrame: | |
| lf = scan_manifest_polars(manifest_path) | |
| cols = lf.collect_schema().names() | |
| token_col = ( | |
| "estimated_token_count" if "estimated_token_count" in cols else "token_count" | |
| ) | |
| return lf.select( | |
| pl.col(token_col).alias("token_count"), | |
| "weborganizer_topic", | |
| "weborganizer_format", | |
| ) | |
| def validate_labels( | |
| lf: pl.LazyFrame, | |
| known_topics: list[str], | |
| known_formats: list[str], | |
| ) -> tuple[list[str], list[str]]: | |
| topics_df = lf.select(pl.col("weborganizer_topic").unique()).collect() | |
| formats_df = lf.select(pl.col("weborganizer_format").unique()).collect() | |
| found_topics = set(topics_df["weborganizer_topic"].to_list()) | |
| found_formats = set(formats_df["weborganizer_format"].to_list()) | |
| unknown_topics = sorted(found_topics - set(known_topics) - {None}) | |
| unknown_formats = sorted(found_formats - set(known_formats) - {None}) | |
| return unknown_topics, unknown_formats | |
| def compute_bin_stats( | |
| lf: pl.LazyFrame, | |
| topics: list[str], | |
| formats: list[str], | |
| target_docs_per_bin: int = TARGET_DOCS_PER_BIN, | |
| ) -> pl.DataFrame: | |
| agg = ( | |
| lf.group_by("weborganizer_topic", "weborganizer_format") | |
| .agg( | |
| pl.len().alias("doc_count"), | |
| pl.col("token_count").sum().alias("token_count"), | |
| pl.col("token_count").mean().alias("mean_tokens_per_doc"), | |
| pl.col("token_count").median().alias("median_tokens_per_doc"), | |
| pl.col("token_count").min().alias("min_tokens_per_doc"), | |
| pl.col("token_count").max().alias("max_tokens_per_doc"), | |
| pl.col("token_count").std().alias("std_tokens_per_doc"), | |
| ) | |
| .collect() | |
| ) | |
| canonical = pl.DataFrame( | |
| { | |
| "weborganizer_topic": [t for t in topics for _ in formats], | |
| "weborganizer_format": formats * len(topics), | |
| } | |
| ) | |
| stats = canonical.join( | |
| agg, | |
| on=["weborganizer_topic", "weborganizer_format"], | |
| how="left", | |
| ).with_columns( | |
| pl.col("doc_count").fill_null(0).cast(pl.UInt64), | |
| pl.col("token_count").fill_null(0), | |
| pl.col("mean_tokens_per_doc").fill_null(0.0), | |
| pl.col("median_tokens_per_doc").fill_null(0.0), | |
| pl.col("min_tokens_per_doc").fill_null(0), | |
| pl.col("max_tokens_per_doc").fill_null(0), | |
| pl.col("std_tokens_per_doc").fill_null(0.0), | |
| ) | |
| stats = stats.with_columns( | |
| (pl.col("weborganizer_topic") + "__" + pl.col("weborganizer_format")).alias( | |
| "bin_id" | |
| ), | |
| pl.when(pl.col("doc_count") == 0) | |
| .then(pl.lit(BIN_EMPTY)) | |
| .when(pl.col("doc_count") < 1_000) | |
| .then(pl.lit(BIN_SPARSE)) | |
| .when(pl.col("doc_count") < target_docs_per_bin) | |
| .then(pl.lit(BIN_PARTIAL)) | |
| .otherwise(pl.lit(BIN_FULL)) | |
| .alias("classification"), | |
| ) | |
| return stats | |
| def detect_length_mismatches( | |
| bin_stats: pl.DataFrame, | |
| doc_threshold: int = DEFAULT_DOC_THRESHOLD, | |
| token_floor: int = DEFAULT_TOKEN_FLOOR, | |
| token_cap: int = DEFAULT_TOKEN_CAP, | |
| ) -> pl.DataFrame: | |
| eligible = bin_stats.filter(pl.col("doc_count") >= doc_threshold) | |
| if eligible.is_empty(): | |
| return eligible.with_columns(pl.lit("").alias("flags")) | |
| flag_exprs = { | |
| "short_docs": pl.col("token_count") < token_floor, | |
| "long_docs": pl.col("token_count") > token_cap, | |
| "high_variance": ( | |
| (pl.col("median_tokens_per_doc") > 0) | |
| & (pl.col("max_tokens_per_doc") / pl.col("median_tokens_per_doc") > 100) | |
| ), | |
| "heavy_skew": ( | |
| (pl.col("median_tokens_per_doc") > 0) | |
| & (pl.col("mean_tokens_per_doc") > 3 * pl.col("median_tokens_per_doc")) | |
| ), | |
| } | |
| flagged = eligible.with_columns( | |
| **{name: expr.alias(name) for name, expr in flag_exprs.items()} | |
| ) | |
| names = list(flag_exprs) | |
| any_flag = pl.any_horizontal(*[pl.col(n) for n in names]) | |
| flagged = ( | |
| flagged.filter(any_flag) | |
| .with_columns( | |
| pl.concat_str( | |
| [ | |
| pl.when(pl.col(c)).then(pl.lit(c)).otherwise(pl.lit("")) | |
| for c in names | |
| ], | |
| separator=",", | |
| ) | |
| .str.replace_all(r",{2,}", ",") | |
| .str.strip_chars(",") | |
| .alias("flags"), | |
| ) | |
| .drop(names) | |
| ) | |
| return flagged | |
| __all__ = [ | |
| "DEFAULT_DOC_THRESHOLD", | |
| "DEFAULT_TOKEN_CAP", | |
| "DEFAULT_TOKEN_FLOOR", | |
| "compute_bin_stats", | |
| "detect_length_mismatches", | |
| "load_manifest", | |
| "validate_labels", | |
| ] | |
Xet Storage Details
- Size:
- 5.12 kB
- Xet hash:
- 0d67c1d4d9f1a6c868bd77d1a39b617a9040e2b59391c69502a1ad8b55866433
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.