Claude
Skills
Sign in
Back

tensorflow-model-deployment

Included with Lifetime
$97 forever

Deploy and serve TensorFlow models

Cloud & DevOps

What this skill does


# TensorFlow Model Deployment

Deploy TensorFlow models to production environments using SavedModel format, TensorFlow Lite for mobile and edge devices, quantization techniques, and serving infrastructure. This skill covers model export, optimization, conversion, and deployment strategies.

## SavedModel Export

### Basic SavedModel Export

```python
# Save model to TensorFlow SavedModel format
model.save('path/to/saved_model')

# Load SavedModel
loaded_model = tf.keras.models.load_model('path/to/saved_model')

# Make predictions with loaded model
predictions = loaded_model.predict(test_data)
```

### Create Serving Model

```python
# Create serving model from classifier
serving_model = classifier.create_serving_model()

# Inspect model inputs and outputs
print(f'Model\'s input shape and type: {serving_model.inputs}')
print(f'Model\'s output shape and type: {serving_model.outputs}')

# Save serving model
serving_model.save('model_path')
```

### Export with Signatures

```python
# Define serving signature
@tf.function(input_signature=[tf.TensorSpec(shape=[None, 224, 224, 3], dtype=tf.float32)])
def serve(images):
    return model(images, training=False)

# Save with signature
tf.saved_model.save(
    model,
    'saved_model_dir',
    signatures={'serving_default': serve}
)
```

## TensorFlow Lite Conversion

### Basic TFLite Conversion

```python
# Convert SavedModel to TFLite
converter = tf.lite.TFLiteConverter.from_saved_model('saved_model_dir')
tflite_model = converter.convert()

# Save TFLite model
with open('model.tflite', 'wb') as f:
    f.write(tflite_model)
```

### From Keras Model

```python
# Convert Keras model directly to TFLite
converter = tf.lite.TFLiteConverter.from_keras_model(model)
tflite_model = converter.convert()

# Save to file
import pathlib
tflite_models_dir = pathlib.Path("tflite_models/")
tflite_models_dir.mkdir(exist_ok=True, parents=True)

tflite_model_file = tflite_models_dir / "mnist_model.tflite"
tflite_model_file.write_bytes(tflite_model)
```

### From Concrete Functions

```python
# Convert from concrete function
concrete_function = model.signatures['serving_default']

converter = tf.lite.TFLiteConverter.from_concrete_functions(
    [concrete_function]
)
tflite_model = converter.convert()
```

### Export with Model Maker

```python
# Export trained model to TFLite with metadata
model.export(
    export_dir='output/',
    tflite_filename='model.tflite',
    label_filename='labels.txt',
    vocab_filename='vocab.txt'
)

# Export multiple formats
model.export(
    export_dir='output/',
    export_format=[
        mm.ExportFormat.TFLITE,
        mm.ExportFormat.SAVED_MODEL,
        mm.ExportFormat.LABEL
    ]
)
```

## Model Quantization

### Post-Training Float16 Quantization

```python
from tflite_model_maker.config import QuantizationConfig

# Create float16 quantization config
config = QuantizationConfig.for_float16()

# Export with quantization
model.export(
    export_dir='.',
    tflite_filename='model_fp16.tflite',
    quantization_config=config
)
```

### Dynamic Range Quantization

```python
# Convert with dynamic range quantization
converter = tf.lite.TFLiteConverter.from_saved_model('saved_model_dir')
converter.optimizations = [tf.lite.Optimize.DEFAULT]
tflite_model = converter.convert()

# Save quantized model
with open('model_quantized.tflite', 'wb') as f:
    f.write(tflite_model)
```

### Full Integer Quantization

```python
def representative_dataset():
    """Generate representative dataset for calibration."""
    for i in range(100):
        yield [np.random.rand(1, 224, 224, 3).astype(np.float32)]

# Convert with full integer quantization
converter = tf.lite.TFLiteConverter.from_saved_model('saved_model_dir')
converter.optimizations = [tf.lite.Optimize.DEFAULT]
converter.representative_dataset = representative_dataset
converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8]
converter.inference_input_type = tf.int8
converter.inference_output_type = tf.int8

tflite_model = converter.convert()
```

### Debug Quantization

```python
from tensorflow.lite.python import convert

# Create debug model with numeric verification
converter.optimizations = [tf.lite.Optimize.DEFAULT]
converter.representative_dataset = calibration_gen
converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8]

# Calibrate and quantize with verification
converter._experimental_calibrate_only = True
calibrated = converter.convert()
debug_model = convert.mlir_quantize(calibrated, enable_numeric_verify=True)
```

### Get Quantization Converter

```python
# Apply quantization settings to converter
def get_converter_with_quantization(converter, **kwargs):
    """Apply quantization configuration to converter."""
    config = QuantizationConfig(**kwargs)
    return config.get_converter_with_quantization(converter)

# Use with custom settings
converter = tf.lite.TFLiteConverter.from_keras_model(model)
quantized_converter = get_converter_with_quantization(
    converter,
    optimizations=[tf.lite.Optimize.DEFAULT],
    representative_dataset=representative_dataset
)
tflite_model = quantized_converter.convert()
```

## JAX to TFLite Conversion

### Basic JAX Conversion

```python
from orbax.export import ExportManager
from orbax.export import JaxModule
from orbax.export import ServingConfig
import tensorflow as tf
import jax.numpy as jnp

def model_fn(_, x):
    return jnp.sin(jnp.cos(x))

jax_module = JaxModule({}, model_fn, input_polymorphic_shape='b, ...')

# Option 1: Direct SavedModel conversion
tf.saved_model.save(
    jax_module,
    '/some/directory',
    signatures=jax_module.methods[JaxModule.DEFAULT_METHOD_KEY].get_concrete_function(
        tf.TensorSpec(shape=(None,), dtype=tf.float32, name="input")
    ),
    options=tf.saved_model.SaveOptions(experimental_custom_gradients=True),
)
converter = tf.lite.TFLiteConverter.from_saved_model('/some/directory')
tflite_model = converter.convert()
```

### JAX with Pre/Post Processing

```python
# Option 2: With preprocessing and postprocessing
serving_config = ServingConfig(
    'Serving_default',
    input_signature=[tf.TensorSpec(shape=(None,), dtype=tf.float32, name='input')],
    tf_preprocessor=lambda x: x,
    tf_postprocessor=lambda out: {'output': out}
)
export_mgr = ExportManager(jax_module, [serving_config])
export_mgr.save('/some/directory')

converter = tf.lite.TFLiteConverter.from_saved_model('/some/directory')
tflite_model = converter.convert()
```

### JAX ResNet50 Example

```python
from orbax.export import ExportManager, JaxModule, ServingConfig

# Wrap the model params and function into a JaxModule
jax_module = JaxModule({}, jax_model.apply, trainable=False)

# Specify the serving configuration and export the model
serving_config = ServingConfig(
    "serving_default",
    input_signature=[tf.TensorSpec([480, 640, 3], tf.float32, name="inputs")],
    tf_preprocessor=resnet_image_processor,
    tf_postprocessor=lambda x: tf.argmax(x, axis=-1),
)

export_manager = ExportManager(jax_module, [serving_config])

saved_model_dir = "resnet50_saved_model"
export_manager.save(saved_model_dir)

# Convert to TFLite
converter = tf.lite.TFLiteConverter.from_saved_model(saved_model_dir)
tflite_model = converter.convert()
```

## Model Optimization

### Graph Transformation

```bash
# Build graph transformation tool
bazel build tensorflow/tools/graph_transforms:transform_graph

# Optimize for deployment
bazel-bin/tensorflow/tools/graph_transforms/transform_graph \
--in_graph=tensorflow_inception_graph.pb \
--out_graph=optimized_inception_graph.pb \
--inputs='Mul' \
--outputs='softmax' \
--transforms='
  strip_unused_nodes(type=float, shape="1,299,299,3")
  remove_nodes(op=Identity, op=CheckNumerics)
  fold_constants(ignore_errors=true)
  fold_batch_norms
  fold_old_batch_norms'
```

### Fix Mobile Kernel Errors

```bash
# Optimize for mobile deployment
bazel-bin/tensorflow/tools/graph_transforms/transform_graph \
--in_graph=tensorflow_inception_gra

Related in Cloud & DevOps