| 24 | model_name = "gpt2" |
| 25 | |
| 26 | def 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( |