Model function.
(features, labels, mode, config)
| 152 | """Creates a model function.""" |
| 153 | |
| 154 | def _model_fn(features, labels, mode, config): |
| 155 | """Model function.""" |
| 156 | assert labels is None, labels |
| 157 | (loss, |
| 158 | scores, |
| 159 | model_predictions, |
| 160 | training_op, |
| 161 | init_op, |
| 162 | is_initialized) = gmm_ops.gmm(self._parse_tensor_or_dict(features), |
| 163 | self._training_initial_clusters, |
| 164 | self._num_clusters, self._random_seed, |
| 165 | self._covariance_type, |
| 166 | self._params) |
| 167 | incr_step = state_ops.assign_add(training_util.get_global_step(), 1) |
| 168 | training_op = with_dependencies([training_op, incr_step], loss) |
| 169 | training_hooks = [_InitializeClustersHook( |
| 170 | init_op, is_initialized, config.is_chief)] |
| 171 | predictions = { |
| 172 | GMM.ASSIGNMENTS: model_predictions[0][0], |
| 173 | } |
| 174 | eval_metric_ops = { |
| 175 | GMM.SCORES: scores, |
| 176 | GMM.LOG_LIKELIHOOD: _streaming_sum(loss), |
| 177 | } |
| 178 | return model_fn_lib.ModelFnOps(mode=mode, predictions=predictions, |
| 179 | eval_metric_ops=eval_metric_ops, |
| 180 | loss=loss, train_op=training_op, |
| 181 | training_hooks=training_hooks) |
| 182 | |
| 183 | return _model_fn |
no test coverage detected