vibrato-v1 / export_coreml.py
rrodriguez50928's picture
Vibrato v1.0 — initial open release (VocalSet-trained multi-task vocal analysis CNN)
253daab verified
Raw
History Blame Contribute Delete
4.15 kB
"""
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)