CoolFace
Modelpublic

Yehor/YAMNet-CoreML

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
0likes28downloads
main.py358 linesDownload Raw Back to root
1import argparse2import shutil3import tarfile4import tempfile5import urllib.request6from dataclasses import dataclass7from pathlib import Path8 9import coremltools as ct10import tensorflow as tf11 12DEFAULT_TFHUB_URL = "https://tfhub.dev/google/yamnet/1?tf-hub-format=compressed"13 14 15def parse_args():16    script_dir = Path(__file__).resolve().parent17    parser = argparse.ArgumentParser(description="Convert YAMNet TensorFlow SavedModel to Core ML.")18    parser.add_argument(19        "--model-path",20        default=script_dir / "yamnet_model",21        type=Path,22        help="Path to a TensorFlow SavedModel directory, .keras file, or .h5 file.",23    )24    parser.add_argument(25        "--output",26        default=script_dir / "YAMNet.mlpackage",27        type=Path,28        help="Output Core ML package path.",29    )30    parser.add_argument(31        "--waveform-samples",32        default=15_600,33        type=int,34        help="Fixed waveform length for the Core ML input. YAMNet commonly uses 0.975s at 16 kHz.",35    )36    parser.add_argument(37        "--output-key",38        default=None,39        help="SavedModel output key for --conversion-mode waveform. Defaults to output_0 for TFHub YAMNet.",40    )41    parser.add_argument(42        "--conversion-mode",43        choices=["features", "waveform"],44        default="features",45        help=(46            "features converts the YAMNet classifier from 96x64 log-mel patches and avoids unsupported "47            "TensorFlow FFT ops. waveform attempts direct full SavedModel conversion."48        ),49    )50    parser.add_argument(51        "--download-tfhub",52        action="store_true",53        help="Download YAMNet from TFHub into --model-path before conversion if it is missing.",54    )55    parser.add_argument(56        "--force-download",57        action="store_true",58        help="Replace --model-path with a fresh TFHub download before conversion.",59    )60    parser.add_argument(61        "--download-only",62        action="store_true",63        help="Download YAMNet from TFHub into --model-path and exit without Core ML conversion.",64    )65    parser.add_argument(66        "--tfhub-url",67        default=DEFAULT_TFHUB_URL,68        help="TFHub compressed SavedModel URL.",69    )70    return parser.parse_args()71 72 73@dataclass(frozen=True)74class YamnetParams:75    patch_frames: int = 9676    patch_bands: int = 6477    num_classes: int = 52178    conv_padding: str = "same"79    batchnorm_center: bool = True80    batchnorm_scale: bool = False81    batchnorm_epsilon: float = 1e-482    classifier_activation: str = "sigmoid"83 84 85YAMNET_LAYER_DEFS = [86    ("conv", [3, 3], 2, 32),87    ("separable_conv", [3, 3], 1, 64),88    ("separable_conv", [3, 3], 2, 128),89    ("separable_conv", [3, 3], 1, 128),90    ("separable_conv", [3, 3], 2, 256),91    ("separable_conv", [3, 3], 1, 256),92    ("separable_conv", [3, 3], 2, 512),93    ("separable_conv", [3, 3], 1, 512),94    ("separable_conv", [3, 3], 1, 512),95    ("separable_conv", [3, 3], 1, 512),96    ("separable_conv", [3, 3], 1, 512),97    ("separable_conv", [3, 3], 1, 512),98    ("separable_conv", [3, 3], 2, 1024),99    ("separable_conv", [3, 3], 1, 1024),100]101 102 103def download_tfhub_saved_model(tfhub_url, model_path, force=False):104    if model_path.exists() and not force:105        print(f"Using existing model: {model_path}")106        return107 108    if model_path.exists():109        if model_path.is_dir():110            shutil.rmtree(model_path)111        else:112            model_path.unlink()113 114    model_path.parent.mkdir(parents=True, exist_ok=True)115    print(f"Downloading TFHub model: {tfhub_url}")116 117    with tempfile.TemporaryDirectory(prefix="yamnet-tfhub-") as temp_dir:118        archive_path = Path(temp_dir) / "model.tar.gz"119        request = urllib.request.Request(tfhub_url, headers={"User-Agent": "langpipe-yamnet-coreml"})120        with urllib.request.urlopen(request) as response, archive_path.open("wb") as archive:121            shutil.copyfileobj(response, archive)122 123        extract_path = Path(temp_dir) / "model"124        extract_path.mkdir()125        with tarfile.open(archive_path, "r:gz") as tar:126            tar.extractall(extract_path, filter="data")127 128        saved_model_pb = next(extract_path.rglob("saved_model.pb"), None)129        if saved_model_pb is None:130            raise ValueError(f"TFHub archive did not contain saved_model.pb: {tfhub_url}")131 132        extracted_model_root = saved_model_pb.parent133        shutil.copytree(extracted_model_root, model_path)134 135    print(f"Saved TFHub model: {model_path}")136 137 138def batch_norm(name, params, layer_input):139    return tf.keras.layers.BatchNormalization(140        name=name,141        center=params.batchnorm_center,142        scale=params.batchnorm_scale,143        epsilon=params.batchnorm_epsilon,144    )(layer_input)145 146 147def build_yamnet_feature_model(params=YamnetParams()):148    features = tf.keras.Input(149        shape=(params.patch_frames, params.patch_bands),150        dtype=tf.float32,151        name="features",152    )153    net = tf.keras.layers.Reshape(154        (params.patch_frames, params.patch_bands, 1),155        name="features_4d",156    )(features)157 158    for layer_index, (layer_type, kernel, stride, filters) in enumerate(YAMNET_LAYER_DEFS, start=1):159        prefix = f"layer{layer_index}"160        if layer_type == "conv":161            net = tf.keras.layers.Conv2D(162                name=f"{prefix}_conv",163                filters=filters,164                kernel_size=kernel,165                strides=stride,166                padding=params.conv_padding,167                use_bias=False,168                activation=None,169            )(net)170            net = batch_norm(f"{prefix}_conv_bn", params, net)171            net = tf.keras.layers.ReLU(name=f"{prefix}_relu")(net)172            continue173 174        net = tf.keras.layers.DepthwiseConv2D(175            name=f"{prefix}_depthwise_conv",176            kernel_size=kernel,177            strides=stride,178            depth_multiplier=1,179            padding=params.conv_padding,180            use_bias=False,181            activation=None,182        )(net)183        net = batch_norm(f"{prefix}_depthwise_conv_bn", params, net)184        net = tf.keras.layers.ReLU(name=f"{prefix}_depthwise_relu")(net)185        net = tf.keras.layers.Conv2D(186            name=f"{prefix}_pointwise_conv",187            filters=filters,188            kernel_size=(1, 1),189            strides=1,190            padding=params.conv_padding,191            use_bias=False,192            activation=None,193        )(net)194        net = batch_norm(f"{prefix}_pointwise_conv_bn", params, net)195        net = tf.keras.layers.ReLU(name=f"{prefix}_pointwise_relu")(net)196 197    embeddings = tf.keras.layers.GlobalAveragePooling2D(name="embeddings")(net)198    logits = tf.keras.layers.Dense(units=params.num_classes, use_bias=True, name="dense")(embeddings)199    class_scores = tf.keras.layers.Activation(200        activation=params.classifier_activation,201        name="class_scores",202    )(logits)203    return tf.keras.Model(name="yamnet_features", inputs=features, outputs=class_scores)204 205 206def load_feature_model_weights_from_saved_model(model, model_path):207    loaded = tf.saved_model.load(str(model_path))208    if not hasattr(loaded, "_yamnet"):209        raise ValueError("SavedModel does not expose the expected TFHub YAMNet _yamnet object.")210 211    source_weights = list(loaded._yamnet.variables)212    target_weights = model.weights213    if len(source_weights) != len(target_weights):214        raise ValueError(f"Weight count mismatch: source={len(source_weights)}, target={len(target_weights)}")215 216    for index, (target, source) in enumerate(zip(target_weights, source_weights)):217        if tuple(target.shape) != tuple(source.shape):218            raise ValueError(219                f"Weight shape mismatch at {index}: target {target.name} {target.shape}, "220                f"source {source.name} {source.shape}"221            )222 223    model.set_weights([weight.numpy() for weight in source_weights])224 225 226def describe_signature(signature):227    _, keyword_specs = signature.structured_input_signature228    print("SavedModel inputs:")229    for name, spec in keyword_specs.items():230        print(f"  {name}: shape={spec.shape}, dtype={spec.dtype.name}")231 232    print("SavedModel outputs:")233    for name, spec in signature.structured_outputs.items():234        print(f"  {name}: shape={spec.shape}, dtype={spec.dtype.name}")235 236 237def select_output_key(signature, requested_key):238    outputs = signature.structured_outputs239    if requested_key:240        if requested_key not in outputs:241            raise ValueError(f"--output-key {requested_key!r} not found. Available keys: {list(outputs)}")242        return requested_key243 244    if "output_0" in outputs:245        return "output_0"246 247    for key in outputs:248        if "score" in key.lower() or "class" in key.lower():249            return key250 251    if not outputs:252        raise ValueError("SavedModel serving_default has no outputs.")253    return next(iter(outputs))254 255 256class SavedModelClassScores(tf.Module):257    def __init__(self, signature, input_key, output_key):258        super().__init__()259        self.signature = signature260        self.input_key = input_key261        self.output_key = output_key262 263    @tf.function264    def __call__(self, waveform):265        outputs = self.signature(**{self.input_key: waveform})266        return {"class_scores": outputs[self.output_key]}267 268 269def convert_saved_model(model_path, output_path, waveform_samples, output_key):270    loaded = tf.saved_model.load(str(model_path))271    if "serving_default" not in loaded.signatures:272        raise ValueError(f"SavedModel has no serving_default signature. Available: {list(loaded.signatures)}")273 274    signature = loaded.signatures["serving_default"]275    describe_signature(signature)276 277    _, keyword_specs = signature.structured_input_signature278    if len(keyword_specs) != 1:279        raise ValueError(280            "Expected one SavedModel input. Pass a wrapper model if this SavedModel has "281            f"{len(keyword_specs)} inputs: {list(keyword_specs)}"282        )283 284    input_key = next(iter(keyword_specs))285    selected_output = select_output_key(signature, output_key)286    print(f"Converting input {input_key!r} -> output {selected_output!r} as 'class_scores'")287 288    wrapper = SavedModelClassScores(signature, input_key, selected_output)289    concrete = wrapper.__call__.get_concrete_function(290        tf.TensorSpec([waveform_samples], tf.float32, name="waveform")291    )292 293    return ct.convert(294        [concrete],295        source="tensorflow",296        inputs=[ct.TensorType(shape=(waveform_samples,), name="waveform")],297        convert_to="mlprogram",298        minimum_deployment_target=ct.target.macOS13,299    )300 301 302def convert_keras_model(model_path, output_path):303    model = tf.keras.models.load_model(str(model_path))304    return ct.convert(305        model,306        inputs=[ct.TensorType(shape=(1, 64, 96, 1), name="features")],307        outputs=[ct.TensorType(name="class_scores")],308        convert_to="mlprogram",309        minimum_deployment_target=ct.target.macOS13,310    )311 312 313def convert_feature_model(model_path):314    model = build_yamnet_feature_model()315    model(tf.zeros((1, 96, 64), dtype=tf.float32))316    load_feature_model_weights_from_saved_model(model, model_path)317    return ct.convert(318        model,319        source="tensorflow",320        inputs=[ct.TensorType(shape=(1, 96, 64), name="features")],321        convert_to="mlprogram",322        minimum_deployment_target=ct.target.macOS13,323    )324 325 326def main():327    args = parse_args()328    if args.download_tfhub or args.force_download or args.download_only:329        download_tfhub_saved_model(args.tfhub_url, args.model_path, force=args.force_download)330    if args.download_only:331        return332 333    if not args.model_path.exists():334        raise FileNotFoundError(335            f"Model path does not exist: {args.model_path}. "336            "Pass --download-tfhub to download YAMNet from TFHub."337        )338 339    if args.conversion_mode == "features":340        mlmodel = convert_feature_model(args.model_path)341    elif args.model_path.suffix in {".keras", ".h5", ".hdf5"}:342        mlmodel = convert_keras_model(args.model_path, args.output)343    else:344        mlmodel = convert_saved_model(345            args.model_path,346            args.output,347            args.waveform_samples,348            args.output_key,349        )350 351    args.output.parent.mkdir(parents=True, exist_ok=True)352    mlmodel.save(str(args.output))353    print(f"Saved Core ML model: {args.output}")354 355 356if __name__ == "__main__":357    main()358