| 60 | |
| 61 | |
| 62 | def parse_args() -> argparse.Namespace: |
| 63 | parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) |
| 64 | parser.add_argument("--dit_path", required=True, help="Path to the converted Cola DiT directory.") |
| 65 | parser.add_argument("--vae_path", required=True, help="Path to the converted Cola VAE directory.") |
| 66 | parser.add_argument("--tokenizer_path", required=True, help="Path to tokenizer.json.") |
| 67 | parser.add_argument("--max_new_tokens", type=int, default=32) |
| 68 | parser.add_argument("--timestep_num", type=int, default=16) |
| 69 | parser.add_argument("--guidance_scale", type=float, default=7.0) |
| 70 | parser.add_argument("--temperature", type=float, default=0.0) |
| 71 | parser.add_argument("--pad_token_id", type=int, default=100277) |
| 72 | parser.add_argument("--eos_token_id", type=int, default=100257) |
| 73 | parser.add_argument("--im_end_token_id", type=int, default=100265) |
| 74 | return parser.parse_args() |
| 75 | |
| 76 | |
| 77 | def main() -> int: |