()
| 255 | |
| 256 | |
| 257 | def main(): |
| 258 | args = parse_args() |
| 259 | |
| 260 | if args.verb == "dump": |
| 261 | pattern = parse_pattern(args.pattern) |
| 262 | input_length = sum(n for _, n in pattern) |
| 263 | input_words = generate_input_prompt(input_length) |
| 264 | if args.file is not None: |
| 265 | with args.file.open("r") as f: |
| 266 | input_words = f.read().strip().split(" ") |
| 267 | if input_length < sum(n for _, n in pattern): |
| 268 | raise ValueError( |
| 269 | f"Input file has only {input_length} words, but pattern requires at least {input_length} words." |
| 270 | ) |
| 271 | input_length = len(input_words) |
| 272 | logger.info(f"Using {input_length} words") |
| 273 | dump_logits(args.endpoint, args.output, input_words, pattern, args.api_key) |
| 274 | elif args.verb == "compare": |
| 275 | compare_logits(args.input1, args.input2, args.output) |
| 276 | else: |
| 277 | raise ValueError(f"Unknown verb: {args.verb}") |
| 278 | |
| 279 | |
| 280 | if __name__ == "__main__": |
no test coverage detected