MCPcopy Create free account
hub / github.com/PAIR-code/lit / _get_preds

Method _get_preds

lit_nlp/app.py:255–315  ·  view source on GitHub ↗

Get model predictions. Args: data: data payload, containing 'inputs' field model: name of the model to run requested_types: optional, comma-separated list of type names to return requested_fields: optional, comma-separated list of field names to return in additio

(self,
                 data: types.JsonDict,
                 model: Optional[str] = None,
                 requested_types: Optional[str] = None,
                 requested_fields: Optional[str] = None,
                 **kw)

Source from the content-addressed store, hash-verified

253 return dataset.indexed_examples
254
255 def _get_preds(self,
256 data: types.JsonDict,
257 model: Optional[str] = None,
258 requested_types: Optional[str] = None,
259 requested_fields: Optional[str] = None,
260 **kw):
261 """Get model predictions.
262
263 Args:
264 data: data payload, containing 'inputs' field
265 model: name of the model to run
266 requested_types: optional, comma-separated list of type names to return
267 requested_fields: optional, comma-separated list of field names to return
268 in addition to the ones returned due to 'requested_types'.
269 **kw: additional args passed to model.predict()
270
271 Returns:
272 list[JsonDict] containing requested fields of model predictions
273
274 Raises:
275 KeyError: If `data` does not have an 'inputs' property.
276 TypeError: If one of entries in `requested_types` is not a valid LitType.
277 ValueError: If the model returns a different number of predictions than
278 the number of inputs.
279 """
280 if model is None:
281 raise ValueError('Must provide a "model" name to get preds from.')
282
283 inputs = data['inputs']
284 preds = list(self._models[model].predict(
285 [ex['data'] for ex in inputs], **kw))
286
287 num_preds = len(preds)
288 num_inputs = len(inputs)
289 if num_preds != num_inputs:
290 raise ValueError(
291 f'Different number of model predictions ({num_preds}) than inputs'
292 f' ({num_inputs}).'
293 )
294
295 if not requested_types and not requested_fields:
296 return preds
297
298 # Figure out what to return to the frontend.
299 output_spec = self._get_model_spec(model)['output']
300 requested_types = requested_types.split(',') if requested_types else []
301 requested_fields = requested_fields.split(',') if requested_fields else []
302 logging.info('Requested types: %s, fields: %s', str(requested_types),
303 str(requested_fields))
304 for t_name in requested_types:
305 t_class = getattr(types, t_name, None)
306 if not issubclass(t_class, types.LitType):
307 raise TypeError(f"Class '{t_name}' is not a valid LitType.")
308 requested_fields.extend(utils.find_spec_keys(output_spec, t_class))
309 ret_keys = set(requested_fields) # de-dupe
310
311 # Return selected keys.
312 logging.info('Will return keys: %s', str(ret_keys))

Callers 2

_get_metricsMethod · 0.95
_warm_startMethod · 0.95

Calls 3

_get_model_specMethod · 0.95
infoMethod · 0.80
predictMethod · 0.45

Tested by

no test coverage detected