MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / __call__

Method __call__

tensorflow/contrib/compiler/xla.py:72–116  ·  view source on GitHub ↗
(self, features, labels, mode, params)

Source from the content-addressed store, hash-verified

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()."""

Callers

nothing calls this directly

Calls 7

_call_model_fnMethod · 0.95
_make_train_stepMethod · 0.95
_make_eval_stepMethod · 0.95
compileFunction · 0.85
_get_scaffoldFunction · 0.70
getMethod · 0.45
identityMethod · 0.45

Tested by

no test coverage detected