Buckets:

glennmatlin's picture
download
raw
5.12 kB
"""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.