Get value in a specific shape. For each dim, if the required shape is smaller than current shape, ndarray will be sliced. Otherwise, it will be padded with padding_constant at the end. Args: key (_KT): Key in dict. The value of this key must be
(self,
key: _KT,
shape: Union[list, tuple],
padding_constant: int = 0)
| 481 | return value |
| 482 | |
| 483 | def get_value_in_shape(self, |
| 484 | key: _KT, |
| 485 | shape: Union[list, tuple], |
| 486 | padding_constant: int = 0) -> np.ndarray: |
| 487 | """Get value in a specific shape. For each dim, if the required shape |
| 488 | is smaller than current shape, ndarray will be sliced. Otherwise, it |
| 489 | will be padded with padding_constant at the end. |
| 490 | |
| 491 | Args: |
| 492 | key (_KT): |
| 493 | Key in dict. The value of this key must be |
| 494 | an instance of numpy.ndarray. |
| 495 | shape (Union[list, tuple]): |
| 496 | Shape of the returned array. Its length |
| 497 | must be equal to value.ndim. Set -1 for |
| 498 | a dimension if you do not want to edit it. |
| 499 | padding_constant (int, optional): |
| 500 | The value to set the padded values for each axis. |
| 501 | Defaults to 0. |
| 502 | |
| 503 | Raises: |
| 504 | ValueError: |
| 505 | A value in shape is neither positive integer nor -1. |
| 506 | |
| 507 | Returns: |
| 508 | np.ndarray: |
| 509 | An array in required shape. |
| 510 | """ |
| 511 | value = self.get_raw_value(key) |
| 512 | assert isinstance(value, np.ndarray) |
| 513 | assert value.ndim == len(shape) |
| 514 | pad_width_list = [] |
| 515 | slice_list = [] |
| 516 | for dim_index in range(len(shape)): |
| 517 | if shape[dim_index] == -1: |
| 518 | # no pad or slice |
| 519 | pad_width_list.append((0, 0)) |
| 520 | slice_list.append(slice(None)) |
| 521 | elif shape[dim_index] > 0: |
| 522 | # valid shape value |
| 523 | wid = shape[dim_index] - value.shape[dim_index] |
| 524 | if wid > 0: |
| 525 | pad_width_list.append((0, wid)) |
| 526 | else: |
| 527 | pad_width_list.append((0, 0)) |
| 528 | slice_list.append(slice(0, shape[dim_index])) |
| 529 | else: |
| 530 | # invalid |
| 531 | raise ValueError |
| 532 | pad_value = np.pad(value, |
| 533 | pad_width=pad_width_list, |
| 534 | mode='constant', |
| 535 | constant_values=padding_constant) |
| 536 | return pad_value[tuple(slice_list)] |
| 537 | |
| 538 | @overload |
| 539 | def get_slice(self, stop: int): |
nothing calls this directly
no test coverage detected