| 82 | |
| 83 | |
| 84 | def parse_args() -> argparse.Namespace: |
| 85 | parser = argparse.ArgumentParser( |
| 86 | description="Supervised finetuning for MossTTSDelay-family tasks." |
| 87 | ) |
| 88 | parser.add_argument("--model-path", type=str, default="OpenMOSS-Team/MOSS-TTS-v1.5") |
| 89 | parser.add_argument("--codec-path", type=str, default="OpenMOSS-Team/MOSS-Audio-Tokenizer") |
| 90 | parser.add_argument( |
| 91 | "--train-jsonl", |
| 92 | type=str, |
| 93 | required=True, |
| 94 | help="Supports a single JSONL, a directory, a glob, or a comma-separated list of JSONL files.", |
| 95 | ) |
| 96 | parser.add_argument("--output-dir", type=str, default="output") |
| 97 | parser.add_argument("--per-device-batch-size", type=int, default=1) |
| 98 | parser.add_argument("--gradient-accumulation-steps", type=int, default=1) |
| 99 | parser.add_argument("--learning-rate", type=float, default=1e-5) |
| 100 | parser.add_argument("--weight-decay", type=float, default=0.1) |
| 101 | parser.add_argument("--adam-beta1", type=float, default=0.9) |
| 102 | parser.add_argument("--adam-beta2", type=float, default=0.95) |
| 103 | parser.add_argument("--adam-eps", type=float, default=1e-8) |
| 104 | parser.add_argument("--warmup-steps", type=int, default=0) |
| 105 | parser.add_argument("--warmup-ratio", type=float, default=0.03) |
| 106 | parser.add_argument("--lr-scheduler-type", type=str, default="linear", choices=SCHEDULER_CHOICES) |
| 107 | parser.add_argument("--num-epochs", type=int, default=3) |
| 108 | parser.add_argument("--max-train-steps", type=int, default=None) |
| 109 | parser.add_argument("--max-grad-norm", type=float, default=1.0) |
| 110 | parser.add_argument("--logging-steps", type=int, default=1) |
| 111 | parser.add_argument( |
| 112 | "--wandb-project", |
| 113 | type=str, |
| 114 | default=None, |
| 115 | help="If set, log metrics to Weights & Biases (main process only). Requires: pip install wandb", |
| 116 | ) |
| 117 | parser.add_argument("--wandb-run-name", type=str, default=None) |
| 118 | parser.add_argument("--wandb-entity", type=str, default=None) |
| 119 | parser.add_argument( |
| 120 | "--wandb-tags", |
| 121 | type=str, |
| 122 | default=None, |
| 123 | help="Comma-separated tags for the W&B run.", |
| 124 | ) |
| 125 | parser.add_argument("--num-workers", type=int, default=0) |
| 126 | parser.add_argument("--mixed-precision", type=str, default="bf16", choices=["no", "fp16", "bf16"]) |
| 127 | parser.add_argument("--attn-implementation", type=str, default="auto") |
| 128 | parser.add_argument("--audio-tokenizer-device", type=str, default=None) |
| 129 | parser.add_argument("--n-vq", type=int, default=None) |
| 130 | parser.add_argument("--gradient-checkpointing", action="store_true") |
| 131 | parser.add_argument( |
| 132 | "--channelwise-loss-weight", |
| 133 | type=str, |
| 134 | default="1,32", |
| 135 | help=( |
| 136 | "Comma-separated loss weights. Use either n_vq+1 values " |
| 137 | "(text_head,vq0,...,vqN) or two values " |
| 138 | "(text_weight,total_audio_weight). When two values are given, " |
| 139 | "the total audio weight is evenly distributed across all audio heads." |
| 140 | ), |
| 141 | ) |