(self, features, labels, mode, params)
| 70 | self._model_fn = function |
| 71 | |
| 72 | def __call__(self, features, labels, mode, params): |
| 73 | |
| 74 | # TPUEstimator compiles model_fn when use_tpu=True. To avoid double |
| 75 | # compilation, we use this params['use_tpu'] as a hint. When it is set to |
| 76 | # True, model_fn is called without compilation. |
| 77 | # Note that this condition isn't accurate for the case of exporting a model. |
| 78 | # In that case we should ideally not compile so that user can see detailed |
| 79 | # graph. However, we don't have enough information to tell whether model_fn |
| 80 | # is being called for export mode or not. |
| 81 | # TODO(ycao): Make this condition more accurate when implementing PREDICT |
| 82 | # mode. |
| 83 | if params.get('use_tpu'): |
| 84 | return self._call_model_fn(features, labels, mode, params) |
| 85 | |
| 86 | if mode == model_fn_lib.ModeKeys.TRAIN: |
| 87 | train_step, captured_scaffold_fn = self._make_train_step( |
| 88 | features, labels, params) |
| 89 | (loss,) = compile(train_step) |
| 90 | return model_fn_lib.EstimatorSpec( |
| 91 | mode=mode, |
| 92 | loss=loss, |
| 93 | train_op=array_ops.identity(loss), |
| 94 | scaffold=_get_scaffold(captured_scaffold_fn)) |
| 95 | elif mode == model_fn_lib.ModeKeys.EVAL: |
| 96 | eval_step, captured_eval_metric_fn, captured_scaffold_fn = ( |
| 97 | self._make_eval_step(features, labels, params)) |
| 98 | outputs = compile(eval_step) |
| 99 | loss = outputs[0] |
| 100 | |
| 101 | # Calculate eval_metric_ops if eval_metric_fn is set and captured. |
| 102 | eval_metric_fn = captured_eval_metric_fn.get() |
| 103 | if eval_metric_fn: |
| 104 | eval_metric_fn_tensors = outputs[1:] |
| 105 | eval_metric_ops = eval_metric_fn(*eval_metric_fn_tensors) |
| 106 | else: |
| 107 | eval_metric_ops = None |
| 108 | |
| 109 | return model_fn_lib.EstimatorSpec( |
| 110 | mode=mode, |
| 111 | loss=loss, |
| 112 | eval_metric_ops=eval_metric_ops, |
| 113 | scaffold=_get_scaffold(captured_scaffold_fn)) |
| 114 | else: |
| 115 | raise NotImplementedError('%s is not implemented, only TRAIN and EVAL are' |
| 116 | ' supported' % mode) |
| 117 | |
| 118 | def _make_train_step(self, features, labels, params): |
| 119 | """Creates a single step of training for xla.compile().""" |
nothing calls this directly
no test coverage detected