Yehor/YAMNet-CoreML
028
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 