()
| 14 | |
| 15 | |
| 16 | def main(): |
| 17 | parser = argparse.ArgumentParser() |
| 18 | parser.add_argument("input") |
| 19 | parser.add_argument("sample_output", help="train output file") |
| 20 | parser.add_argument("remainder_output", help="valid output file") |
| 21 | parser.add_argument("-k", type=int, help="remainder size") |
| 22 | parser.add_argument( |
| 23 | "--lines", action="store_true", help="split lines instead of docs" |
| 24 | ) |
| 25 | args = parser.parse_args() |
| 26 | |
| 27 | assert args.k is not None |
| 28 | |
| 29 | sample = [] |
| 30 | remainder = [] |
| 31 | num_docs = [0] |
| 32 | |
| 33 | def update_sample(doc): |
| 34 | if len(sample) < args.k: |
| 35 | sample.append(doc.copy()) |
| 36 | else: |
| 37 | i = num_docs[0] |
| 38 | j = random.randrange(i + 1) |
| 39 | if j < args.k: |
| 40 | remainder.append(sample[j]) |
| 41 | sample[j] = doc.copy() |
| 42 | else: |
| 43 | remainder.append(doc.copy()) |
| 44 | num_docs[0] += 1 |
| 45 | doc.clear() |
| 46 | |
| 47 | with open(args.input, "r", encoding="utf-8") as h: |
| 48 | doc = [] |
| 49 | for i, line in enumerate(h): |
| 50 | if line.strip() == "": # empty line indicates new document |
| 51 | update_sample(doc) |
| 52 | else: |
| 53 | doc.append(line) |
| 54 | if args.lines: |
| 55 | update_sample(doc) |
| 56 | if i % 1000000 == 0: |
| 57 | print(i, file=sys.stderr, end="", flush=True) |
| 58 | elif i % 100000 == 0: |
| 59 | print(".", file=sys.stderr, end="", flush=True) |
| 60 | if len(doc) > 0: |
| 61 | update_sample(doc) |
| 62 | print(file=sys.stderr, flush=True) |
| 63 | |
| 64 | assert len(sample) == args.k |
| 65 | |
| 66 | with open(args.sample_output, "w", encoding="utf-8") as out: |
| 67 | first = True |
| 68 | for doc in sample: |
| 69 | if not first and not args.lines: |
| 70 | out.write("\n") |
| 71 | first = False |
| 72 | for line in doc: |
| 73 | out.write(line) |
no test coverage detected