Vibrato v1.0 — initial open release (VocalSet-trained multi-task vocal analysis CNN)
253daab verified | """ | |
| Export trained PyTorch model to CoreML format for iOS deployment. | |
| Usage: | |
| python export_coreml.py --checkpoint checkpoints/best.pt --output ToneyMultiTaskV2.mlpackage | |
| The exported model can be dropped into: | |
| Packages/ToneyCore/Sources/ToneyCore/Resources/ | |
| """ | |
| import argparse | |
| import torch | |
| import coremltools as ct | |
| from model import get_model, VOICE_TYPES, TECHNIQUES, VOWELS, QUALITY_DIMS | |
| def parse_args(): | |
| parser = argparse.ArgumentParser(description="Export ToneyMultiTask to CoreML") | |
| parser.add_argument("--checkpoint", type=str, required=True, help="Path to PyTorch checkpoint (.pt)") | |
| parser.add_argument("--output", type=str, default="ToneyMultiTaskV2.mlpackage", help="Output .mlpackage path") | |
| parser.add_argument("--quantize", action="store_true", help="Apply float16 quantization") | |
| return parser.parse_args() | |
| class ToneyMultiTaskExport(torch.nn.Module): | |
| """Wrapper that returns a tuple instead of dict (required for tracing).""" | |
| def __init__(self, model): | |
| super().__init__() | |
| self.model = model | |
| def forward(self, audio_input: torch.Tensor): | |
| out = self.model(audio_input) | |
| return ( | |
| out["voice_logits"], | |
| out["technique_logits"], | |
| out["vowel_logits"], | |
| out["quality_scores"], | |
| ) | |
| def export(checkpoint_path: str, output_path: str, quantize: bool = False): | |
| # Load model | |
| model = get_model() | |
| checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True) | |
| model.load_state_dict(checkpoint["model_state_dict"]) | |
| model.eval() | |
| print(f"Loaded checkpoint from: {checkpoint_path}") | |
| if "val_metrics" in checkpoint: | |
| metrics = checkpoint["val_metrics"] | |
| print(f" Epoch: {checkpoint.get('epoch', '?')}") | |
| print(f" Val loss: {metrics.get('val_loss', '?'):.4f}") | |
| print(f" Voice acc: {metrics.get('voice_acc', '?'):.1%}") | |
| print(f" Technique acc: {metrics.get('technique_acc', '?'):.1%}") | |
| print(f" Vowel acc: {metrics.get('vowel_acc', '?'):.1%}") | |
| # Wrap model for export | |
| export_model = ToneyMultiTaskExport(model) | |
| export_model.eval() | |
| # Trace with example input | |
| example_input = torch.randn(1, 4000) | |
| traced = torch.jit.trace(export_model, example_input) | |
| # Convert to CoreML | |
| mlmodel = ct.convert( | |
| traced, | |
| inputs=[ | |
| ct.TensorType(name="audio_input", shape=(1, 4000)), | |
| ], | |
| outputs=[ | |
| ct.TensorType(name="voice_logits"), | |
| ct.TensorType(name="technique_logits"), | |
| ct.TensorType(name="vowel_logits"), | |
| ct.TensorType(name="quality_scores"), | |
| ], | |
| minimum_deployment_target=ct.target.iOS17, | |
| ) | |
| # Add metadata | |
| mlmodel.author = "Toney AI (toney-ai)" | |
| mlmodel.short_description = "Multi-task vocal analysis: voice type, singing technique, vowel, and quality scoring" | |
| mlmodel.version = "2.0" | |
| mlmodel.license = "Apache-2.0" | |
| # Add label descriptions | |
| labels_desc = ( | |
| f"voice_logits: {VOICE_TYPES}\n" | |
| f"technique_logits: {TECHNIQUES}\n" | |
| f"vowel_logits: {VOWELS}\n" | |
| f"quality_scores: {QUALITY_DIMS}" | |
| ) | |
| mlmodel.input_description["audio_input"] = "Raw audio samples, 0.25s window at 16kHz" | |
| mlmodel.output_description["voice_logits"] = f"Voice type logits: {VOICE_TYPES}" | |
| mlmodel.output_description["technique_logits"] = f"Technique logits: {TECHNIQUES}" | |
| mlmodel.output_description["vowel_logits"] = f"Vowel logits: {VOWELS}" | |
| mlmodel.output_description["quality_scores"] = f"Quality scores (0-1): {QUALITY_DIMS}" | |
| # Optional quantization | |
| if quantize: | |
| print("Applying float16 quantization...") | |
| mlmodel = ct.models.neural_network.quantization_utils.quantize_weights(mlmodel, nbits=16) | |
| # Save | |
| mlmodel.save(output_path) | |
| print(f"\nCoreML model saved to: {output_path}") | |
| print(f"\nTo use in Toney iOS app:") | |
| print(f" cp {output_path} Packages/ToneyCore/Sources/ToneyCore/Resources/") | |
| if __name__ == "__main__": | |
| args = parse_args() | |
| export(args.checkpoint, args.output, args.quantize) | |