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)
| 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)) |
no test coverage detected