MCPcopy Create free account
hub / github.com/AnswerDotAI/ModernBERT / main

Function main

src/data/source_stats.py:26–131  ·  view source on GitHub ↗
(out_fn, dataset_max_size)

Source from the content-addressed store, hash-verified

24model_name = "gpt2"
25
26def main(out_fn, dataset_max_size):
27 tokenizer = AutoTokenizer.from_pretrained(model_name, fast=True)
28
29 # load all lines in out_fn
30 percentiles_out_path = str(out_fn).replace(".csv", ".jsonl")
31 cached_sources = set()
32 if os.path.exists(percentiles_out_path):
33 with open(str(out_fn).replace(".csv", ".jsonl"), "r") as f:
34 for line in f:
35 cached_sources.add(json.loads(line)["source"])
36
37
38 percentiles = [1, 99] + list(range(0, 101, 5))
39 print(f"Saving source-level token count percentiles to {str(out_fn).replace('.csv', '.jsonl')}")
40 stats = []
41 tokens_for_source = []
42 current_repos_to_do = [item for item in ALL_REPOS if item.split("/")[-1] not in cached_sources]
43 prev_src = current_repos_to_do[0].split("/")[-1]
44 percentiles_out = open(percentiles_out_path, "a")
45
46 for data_dir in tqdm.tqdm(current_repos_to_do):
47 source = data_dir.split("/")[-1]
48 print(f"Processing {source}... with data_dir {data_dir}")
49 if source != prev_src:
50 # add percentiles and reset
51 tokens_np = np.array(tokens_for_source)
52 percentile_stats_all = np.percentile(tokens_np, percentiles)
53 percentile_stats = {
54 "mean": np.mean(tokens_np),
55 "std": np.std(tokens_np),
56 "percentiles": {p: v for p, v in zip(percentiles, percentile_stats_all)}
57 }
58 tokens_for_source = []
59 percentiles_out.write(json.dumps({prev_src: percentile_stats, "source": prev_src}) + "\n")
60 percentiles_out.flush()
61
62 prev_src = source
63
64 with tempfile.TemporaryDirectory() as tmp_cache_dir:
65 remote = f'hf://datasets/orionweller/{source}/'
66 token_lens = []
67 pool = []
68 clean_stale_shared_memory()
69 for idx, instance in tqdm.tqdm(enumerate(StreamingDataset(remote=remote, shuffle=False, split=None, batch_size=1, predownload=dataset_max_size))):
70 pool.append(instance)
71 if idx > dataset_max_size:
72 break
73 if len(pool) > 1000:
74 hf_dataset = Dataset.from_list(pool)
75 try:
76 tokens = hf_dataset.map(
77 lambda row: {"num_tokens": tokenizer(row["text"]), "batched": True},
78 num_proc=NUM_PROC, remove_columns=MDS_COLS_TEXT.keys()
79 )["num_tokens"]
80 except Exception as e:
81 print(f"Error processing {source} at idx {idx}")
82 print(e)
83 tokens = hf_dataset.map(

Callers 1

source_stats.pyFile · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected