(sources, targets, max_sample=None)
| 79 | raise ValueError(f"{dataset_name} not implement yet") |
| 80 | |
| 81 | def extract_random_dataset(sources, targets, max_sample=None): |
| 82 | if max_sample is not None: |
| 83 | if max_sample <= len(sources): |
| 84 | print(f"only use {max_sample} samples") |
| 85 | random_indices = random.sample(range(len(sources)), max_sample) |
| 86 | sources = [sources[i] for i in random_indices] |
| 87 | targets = [targets[i] for i in random_indices] |
| 88 | else: |
| 89 | print("max_sample exceeds the length of the array. Using the entire array.") |
| 90 | sources = sources |
| 91 | targets = targets |
| 92 | else: |
| 93 | print(f"using the whole {len(sources)} samples") |
| 94 | |
| 95 | return sources, targets |
| 96 | |
| 97 | def get_wiki_dataset(max_sample): |
| 98 | wiki_dataset = load_dataset("wikitext", 'wikitext-2-raw-v1', split='train') |
no outgoing calls
no test coverage detected