tensorflow / tensorflow/model-optimization

Pruning: improve custom training loop API

Open
#271 5 comments 1 reaction 1 assignee View on GitHub

@liyunlu0618 is already working on this.

Since Apr 16, 2021.

feature request priority:low technique:pruning
Dominant language
Python
Stars
1.6k
Forks
349
Avg merge
3d 2h
Merged PRs (30d)
1

Description

The recommended path for pruning with a custom training loop is not as simple as it could be.

pruned_model = setup_pruned_model()

loss = tf.keras.losses.categorical_crossentropy
optimizer = keras.optimizers.Adam()

log_dir = tempfile.mkdtemp()

# This is all not boilerplate.
pruned_model.optimizer = optimizer
step_callback = tfmot.sparsity.keras.UpdatePruningStep()
step_callback.set_model(pruned_model)
log_callback = tfmot.sparsity.keras.PruningSummaries(log_dir=log_dir) # optional Tensorboard logging.
log_callback.set_model(pruned_model)

step_callback.on_train_begin()
for _ in range(3):
    # only one batch given batch_size = 20 and input shape.
    step_callback.on_train_batch_begin(batch=unused_arg)
    inp = np.reshape(x_train,
                     [self._BATCH_SIZE, 10])  # original shape: from [10].
    with tf.GradientTape() as tape:
      logits = pruned_model(inp, training=True)
      loss_value = loss(y_train, logits)
      grads = tape.gradient(loss_value, pruned_model.trainable_variables)
      optimizer.apply_gradients(zip(grads, pruned_model.trainable_variables))

    step_callback.on_epoch_end(batch=unused_arg)
    log_callback.on_epoch_end(batch=unused_arg)
...

The set_model and pruned_model.optimizer setting is unusual and could be missed.

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.