Download code/tt_moge/reference/utils3d/numpy/utils.py from changh95/moge-2-p150: direct link, hf CLI and curl.
- Browser
- Download file 21.6 kB
-
https://huggingface.co/changh95/moge-2-p150/resolve/main/code/tt_moge/reference/utils3d/numpy/utils.py
- Command line
-
hf download hf://changh95/moge-2-p150/code/tt_moge/reference/utils3d/numpy/utils.py
-
curl -L -o utils.py https://huggingface.co/changh95/moge-2-p150/resolve/main/code/tt_moge/reference/utils3d/numpy/utils.py
21.6 kB
| import numpy as np | |
| from numpy import ndarray | |
| from typing import * | |
| from numbers import Number, Integral | |
| import warnings | |
| import functools | |
| import math | |
| if TYPE_CHECKING: | |
| from scipy.sparse import csr_array | |
| __all__ = [ | |
| 'sliding_window', | |
| 'pooling', | |
| 'max_pool_2d', | |
| 'lookup', | |
| 'lookup_get', | |
| 'lookup_set', | |
| 'group', | |
| 'csr_matrix_from_dense_indices', | |
| 'reverse_permutation', | |
| 'vector_outer' | |
| ] | |
| def sliding_window( | |
| x: ndarray, | |
| window_size: Union[int, Tuple[int, ...]], | |
| stride: Optional[Union[int, Tuple[int, ...]]] = None, | |
| dilation: Optional[Union[int, Tuple[int, ...]]] = None, | |
| pad_size: Optional[Union[int, Tuple[int, int], Tuple[Tuple[int, int]]]] = None, | |
| pad_mode: str = 'constant', | |
| pad_value: Number = 0, | |
| axis: Optional[Tuple[int,...]] = None | |
| ) -> ndarray: | |
| """ | |
| Get a sliding window of the input array. Window axis(axes) will be appended as the last dimension(s). | |
| This function is a wrapper of `numpy.lib.stride_tricks.sliding_window_view` with additional support for padding and stride. | |
| ## Parameters | |
| - `x` (ndarray): Input array. | |
| - `window_size` (int or Tuple[int,...]): Size of the sliding window. If int | |
| is provided, the same size is used for all specified axes. | |
| - `stride` (Optional[Tuple[int,...]]): Stride between the sliding windows. If None, | |
| no stride is applied. If int is provided, the same stride is used for all specified axes. | |
| - `dilation` (Optional[Tuple[int,...]]): Dilation in each sliding window. If None, | |
| no dilation is applied. If int is provided, the same dilation is used for all specified axes. | |
| - `pad_size` (Optional[Union[int, Tuple[int, int], Tuple[Tuple[int, int]]]]): Size of padding to apply before sliding window. | |
| Corresponding to `axis`. | |
| - General format is `((before_1, after_1), (before_2, after_2), ...)`. | |
| - Shortcut formats: | |
| - `int` -> same padding before and after for all axes; | |
| - `(int, int)` -> same padding before and after for each axis; | |
| - `((int,), (int,) ...)` -> specify padding for each axis, same before and after. | |
| - `pad_mode` (str): Padding mode to use. Refer to `numpy.pad` for more details. | |
| - `pad_value` (Union[int, float]): Value to use for constant padding. Only used | |
| when `pad_mode` is 'constant'. | |
| - `axis` (Optional[Tuple[int,...]]): Axes to apply the sliding window. If None, all axes are used. | |
| ## Returns | |
| - (ndarray): Sliding window of the input array. | |
| - If no padding, the output is a view of the input array with zero copy. | |
| - Otherwise, the output is no longer a view but a copy of the padded array. | |
| """ | |
| # Process axis | |
| if axis is None: | |
| axis = tuple(range(x.ndim)) | |
| if isinstance(axis, Integral): | |
| axis = (axis,) | |
| axis = [axis[i] % x.ndim for i in range(len(axis))] | |
| if isinstance(window_size, Integral): | |
| window_size = (window_size,) * len(axis) | |
| if dilation is not None: | |
| if isinstance(dilation, Integral): | |
| dilation = (dilation,) * len(axis) | |
| if stride is not None: | |
| if isinstance(stride, Integral): | |
| stride = (stride,) * len(axis) | |
| # Pad the input array if needed | |
| if pad_size is not None: | |
| if isinstance(pad_size, Integral): | |
| pad_size = ((pad_size, pad_size),) * len(axis) | |
| elif isinstance(pad_size, tuple) and len(pad_size) == 2 and all(isinstance(p, Integral) for p in pad_size): | |
| pad_size = (pad_size,) * len(axis) | |
| elif isinstance(pad_size, tuple) and all(isinstance(p, tuple) and 1 <= len(p) <= 2 for p in pad_size): | |
| if len(pad_size) == 1: | |
| pad_size = pad_size * len(axis) | |
| else: | |
| assert len(pad_size) == len(axis), f"pad_size {pad_size} must match the number of axes {len(axis)}" | |
| else: | |
| raise ValueError(f"Invalid pad_size {pad_size}") | |
| full_pad = [(0, 0) if i not in axis else pad_size[axis.index(i)] for i in range(x.ndim)] | |
| if pad_mode == 'constant': | |
| x = np.pad(x, full_pad, mode=pad_mode, constant_values=pad_value) | |
| else: | |
| x = np.pad(x, full_pad, mode=pad_mode) | |
| # Apply sliding window | |
| if dilation is None: | |
| x = np.lib.stride_tricks.sliding_window_view(x, window_size, axis=axis) | |
| else: | |
| window_size_dilated = tuple((window_size[i] - 1) * dilation[i] + 1 for i in range(len(window_size))) | |
| x = np.lib.stride_tricks.sliding_window_view(x, window_size_dilated, axis=axis) | |
| # Apply stride if needed | |
| if stride is not None: | |
| stride_slice = tuple(slice(None) if i not in axis else slice(None, None, stride[axis.index(i)]) for i in range(x.ndim - len(axis))) | |
| x = x[stride_slice] | |
| # Apply dilation if needed | |
| if dilation is not None: | |
| dilation_slice = tuple(slice(None, None, dilation[i]) for i in range(len(axis))) | |
| x = x[(..., *dilation_slice)] | |
| return x | |
| def pooling( | |
| x: ndarray, | |
| kernel_size: Union[int, Tuple[int, ...]], | |
| stride: Optional[Union[int, Tuple[int, ...]]] = None, | |
| padding: Optional[Union[int, Tuple[int, int], Tuple[Tuple[int, int]]]] = None, | |
| axis: Optional[Union[int, Tuple[int, ...]]] = None, | |
| mode: Literal['min', 'max', 'sum', 'mean'] = 'max' | |
| ) -> ndarray: | |
| """Compute the pooling of the input array. | |
| NOTE: NaNs will be ignored. | |
| ## Parameters | |
| - `x` (ndarray): Input array. | |
| - `kernel_size` (int or Tuple[int,...]): Size of the pooling window. | |
| - `stride` (Optional[Tuple[int,...]]): Stride of the pooling window. If None, | |
| no stride is applied. If int is provided, the same stride is used for all specified axes. | |
| - `padding` (Optional[Union[int, Tuple[int, int], Tuple[Tuple[int, int]]]]): Size of padding to apply before pooling. | |
| Corresponding to `axis`. | |
| - General format is `((before_1, after_1), (before_2, after_2), ...)`. | |
| - Shortcut formats: | |
| - `int` -> same padding before and after for all axes; | |
| - `(int, int)` -> same padding before and after for each axis; | |
| - `((int,), (int,) ...)` -> specify padding for each axis, same before and after. | |
| - `axis` (Optional[Tuple[int,...]]): Axes to apply the pooling. If None, all axes are used. | |
| - `mode` (str): Pooling mode. One of 'min', 'max', 'sum', 'mean'. | |
| ## Returns | |
| - (ndarray): Pooled array with the same number of dimensions as input array. | |
| """ | |
| if axis is None: | |
| axis = tuple(range(x.ndim)) | |
| if isinstance(axis, Integral): | |
| axis = (axis,) | |
| axis = [axis[i] % x.ndim for i in range(len(axis))] | |
| if isinstance(kernel_size, Integral): | |
| kernel_size = (kernel_size,) * len(axis) | |
| if not isinstance(stride, tuple): | |
| stride = (stride,) * len(axis) | |
| if padding is not None: | |
| if isinstance(padding, Integral): | |
| padding = ((padding, padding),) * len(axis) | |
| elif isinstance(padding, tuple) and len(padding) == 2 and all(isinstance(p, Integral) for p in padding): | |
| padding = (padding,) * len(axis) | |
| elif isinstance(padding, tuple) and all(isinstance(p, tuple) and 1 <= len(p) <= 2 for p in padding): | |
| if len(padding) == 1: | |
| padding = padding * len(axis) | |
| else: | |
| assert len(padding) == len(axis), f"padding {padding} must match the number of axes {len(axis)}" | |
| else: | |
| raise ValueError(f"Invalid padding {padding}") | |
| else: | |
| padding = ((0, 0),) * len(axis) | |
| if mode == 'max': | |
| pad_mode = 'constant' | |
| pad_value = -np.inf if x.dtype.kind == 'f' else np.iinfo(x.dtype).min | |
| pool_fn = np.nanmax | |
| elif mode == 'min': | |
| pad_mode = 'constant' | |
| pad_value = np.inf if x.dtype.kind == 'f' else np.iinfo(x.dtype).max | |
| pool_fn = np.nanmin | |
| elif mode == 'sum': | |
| pad_mode = 'constant' | |
| pad_value = 0 | |
| pool_fn = np.sum | |
| x = np.where(np.isnan(x), 0, x) | |
| elif mode == 'mean': | |
| mask = ~np.isnan(x) | |
| full_pad = [(0, 0) if i not in axis else padding[axis.index(i)] for i in range(x.ndim)] | |
| x = pooling(np.pad(x, full_pad, mode='edge'), kernel_size, stride, axis=axis, mode='sum') | |
| x /= pooling(np.pad(mask, full_pad, mode='edge'), kernel_size, stride, axis=axis, mode='sum') | |
| return x | |
| else: | |
| raise ValueError(f"Invalid pooling mode {mode}. Supported modes are 'min', 'max', 'sum', 'mean'.") | |
| for i in range(len(axis)): | |
| x = pool_fn( | |
| sliding_window(x, kernel_size[i], stride[i], | |
| pad_size=padding[i], pad_mode=pad_mode, pad_value=pad_value, | |
| axis=axis[i]), | |
| axis=-1 | |
| ) | |
| return x | |
| def max_pool_2d(x: ndarray, kernel_size: Union[int, Tuple[int, int]], stride: Union[int, Tuple[int, int]], padding: Union[int, Tuple[int, int]], axis: Tuple[int, int] = (-2, -1)): | |
| if isinstance(kernel_size, Number): | |
| kernel_size = (kernel_size, kernel_size) | |
| if isinstance(stride, Number): | |
| stride = (stride, stride) | |
| if isinstance(padding, Number): | |
| padding = (padding, padding) | |
| axis = tuple(axis) | |
| return pooling(x, kernel_size, stride, padding, axis, 'max') | |
| def lookup(key: ndarray, query: ndarray) -> ndarray: | |
| """Look up `query` in `key` like a dictionary. Useful for COO indexing. | |
| Parameters | |
| ---- | |
| - `key` (ndarray): shape `(num_keys, *key_shape)`, the array to search in | |
| - `query` (ndarray): shape `(..., *key_shape)`, the array to search for. `...` represents any number of batch dimensions. | |
| Returns | |
| ---- | |
| - `indices` (ndarray): shape `(...,)` indices in `key` for each `query`. If a query is not found in key, the corresponding index will be -1. | |
| Notes | |
| ---- | |
| `O((Q + K) * log(Q + K))` complexity, where `Q` is the number of queries and `K` is the number of keys. | |
| """ | |
| assert key.dtype == query.dtype, "Key and query must have the same dtype" | |
| assert key.shape[1:] == query.shape[query.ndim - key.ndim + 1:], f"Key shape {key.shape} and query shape {query.shape} are not compatible." | |
| num_keys, *key_shape = key.shape | |
| query_batch_shape = query.shape[:query.ndim - key.ndim + 1] | |
| key_item_nbytes = math.prod(key_shape) * key.dtype.itemsize | |
| if key.ndim == 1: | |
| # Fast path 1: 1D keys, can directly sort and search | |
| sorted_indices = np.argsort(key) | |
| key_sorted = key[sorted_indices] | |
| result = np.searchsorted(key_sorted, query, side='left') | |
| mask = (result < num_keys) & (key_sorted[result.clip(0, num_keys - 1)] == query) | |
| result = result.astype(np.int64, copy=False) | |
| result[mask] = sorted_indices[result[mask]] | |
| result[~mask] = -1 | |
| return result.reshape(query_batch_shape) | |
| elif key_item_nbytes <= 8: | |
| # Fast path 2: small keys, can view as int64 and sort/search | |
| query_flat = query.reshape(-1, *key_shape) | |
| key_bytes = np.ascontiguousarray(key).view(np.uint8).reshape(num_keys, key_item_nbytes) | |
| query_bytes = np.ascontiguousarray(query_flat).view(np.uint8).reshape(query_flat.shape[0], key_item_nbytes) | |
| if key_item_nbytes < 8: | |
| pad_width = ((0, 0), (0, 8 - key_item_nbytes)) | |
| key_bytes = np.pad(key_bytes, pad_width, mode='constant') | |
| query_bytes = np.pad(query_bytes, pad_width, mode='constant') | |
| key_i64 = key_bytes.view(np.int64).reshape(-1) | |
| query_i64 = query_bytes.view(np.int64).reshape(-1) | |
| sorted_indices = np.argsort(key_i64) | |
| key_sorted = key_i64[sorted_indices] | |
| result = np.searchsorted(key_sorted, query_i64, side='left') | |
| mask = (result < num_keys) & (key_sorted[result.clip(0, num_keys - 1)] == query_i64) | |
| result = result.astype(np.int64, copy=False) | |
| result[mask] = sorted_indices[result[mask]] | |
| result[~mask] = -1 | |
| return result.reshape(query_batch_shape) | |
| else: | |
| query_flat = query.reshape(-1, *key_shape) | |
| _, index, inverse = np.unique( | |
| np.concatenate([key, query_flat], axis=0), | |
| axis=0, | |
| return_index=True, | |
| return_inverse=True | |
| ) | |
| result = index[inverse[num_keys:]] | |
| result[result >= num_keys] = -1 | |
| return result.reshape(query_batch_shape) | |
| def lookup_get(key: ndarray, value: ndarray, get_key: ndarray, default_value: Union[Number, ndarray] = 0) -> ndarray: | |
| """Dictionary-like get for arrays | |
| ## Parameters | |
| - `key` (ndarray): shape `(N, *key_shape)`, the key array of the dictionary to get from | |
| - `value` (ndarray): shape `(N, *value_shape)`, the value array of the dictionary to get from | |
| - `get_key` (ndarray): shape `(..., *key_shape)`, the key array to get for. `...` represents any number of batch dimensions. | |
| - `default_value` (Union[Number, ndarray]): a scalar or an array broadcastable to shape `(..., *value_shape)`. Value to return if a key in `get_key` is not found in `key`. | |
| ## Returns | |
| `get_value` (ndarray): shape `(..., *value_shape)`, result values corresponding to `get_key` | |
| """ | |
| indices = lookup(key, get_key) | |
| if key.shape[0] == 0: | |
| return np.broadcast_to(np.asarray(default_value, dtype=value.dtype), get_key.shape[:get_key.ndim - key.ndim + 1] + value.shape[1:]) | |
| return np.where( | |
| (indices >= 0)[(..., *((None,) * (value.ndim - 1)))], | |
| value[indices.clip(0, key.shape[0] - 1)], | |
| default_value | |
| ) | |
| def lookup_set(key: ndarray, value: ndarray, set_key: ndarray, set_value: ndarray, append: bool = False, inplace: bool = False) -> Tuple[ndarray, ndarray]: | |
| """Dictionary-like set for arrays. | |
| ## Parameters | |
| - `key` (ndarray): shape `(N, *key_shape)`, the key array of the dictionary to set | |
| - `value` (ndarray): shape `(N, *value_shape)`, the value array of the dictionary to set | |
| - `set_key` (ndarray): shape `(M, *key_shape)`, the key array to set for | |
| - `set_value` (ndarray): shape `(M, *value_shape)`, the value array to set as | |
| - `append` (bool): If True, append the (key, value) pairs in (set_key, set_value) that are not in (key, value) to the result. | |
| - `inplace` (bool): If True, modify the input `value` array | |
| ## Returns | |
| - `result_key` (ndarray): shape `(N_new, *value_shape)`. N_new = N + number of new keys added if append is True, else N. | |
| - `result_value (ndarray): shape `(N_new, *value_shape)` | |
| """ | |
| set_indices = lookup(key, set_key) | |
| if inplace: | |
| assert append is False, "Cannot append when inplace is True" | |
| else: | |
| value = value.copy() | |
| hit = np.where(set_indices >= 0) | |
| value[set_indices[hit]] = set_value[hit] | |
| if append: | |
| missing = np.where(set_indices < 0) | |
| key = np.concatenate([key, set_key[missing]], axis=0) | |
| value = np.concatenate([value, set_value[missing]], axis=0) | |
| return key, value | |
| def take_view(a: ndarray, i: Union[int, slice], axis: int = 0) -> ndarray: | |
| """Take a view of the input array at the specified index along the given axis.""" | |
| return a[(slice(None),) * (axis % a.ndim) + (i,)] | |
| def lite_sum(a: ndarray, axis: int = -1) -> ndarray: | |
| """Compute the sum of the input array along the specified small axis. | |
| """ | |
| result_dtype = np.result_type(a.dtype, 0) | |
| if a.shape[axis] == 0: | |
| return np.zeros(a.shape[:axis] + a.shape[axis + 1:], dtype=result_dtype) | |
| elif a.shape[axis] <= 4: # Sweet point for python loop vs einsum | |
| s = take_view(a, 0, axis=axis).astype(result_dtype, copy=True) | |
| for i in range(1, a.shape[axis]): | |
| s += take_view(a, i, axis=axis) | |
| return s | |
| else: # Einsum is faster than np.sum in most cases | |
| return np.einsum('...i->...', np.moveaxis(a, axis, -1), optimize=False) | |
| def lite_prod(a: ndarray, axis: int = -1) -> ndarray: | |
| """Compute the product of the input array along the specified small axis. | |
| """ | |
| result_dtype = np.result_type(a.dtype, 1) | |
| if a.shape[axis] == 0: | |
| return np.ones(a.shape[:axis] + a.shape[axis + 1:], dtype=result_dtype) | |
| elif a.shape[axis] <= 8: | |
| p = take_view(a, 0, axis=axis).astype(result_dtype, copy=True) | |
| for i in range(1, a.shape[axis]): | |
| p *= take_view(a, i, axis=axis) | |
| return p | |
| else: | |
| return np.prod(a, axis=axis) | |
| def lite_dot(a: ndarray, b: ndarray, axis: int = -1) -> ndarray: | |
| """Compute the dot product of two input arrays along the specified small axis. | |
| """ | |
| if a.shape[axis] == 0: | |
| return np.zeros(a.shape[:axis] + a.shape[axis + 1:], dtype=np.result_type(a.dtype, b.dtype)) | |
| elif a.shape[axis] <= 3: | |
| return lite_sum(a * b, axis=axis) | |
| else: | |
| return np.einsum('...i,...i->...', np.moveaxis(a, axis, -1), np.moveaxis(b, axis, -1), optimize=False) | |
| def lite_norm(a: ndarray, ord: int = 2, axis: int = -1) -> ndarray: | |
| """Compute the norm of the input array along the specified small axis. | |
| """ | |
| if ord == 1: | |
| return lite_sum(np.abs(a), axis=axis) | |
| elif ord == 2: | |
| return np.sqrt(lite_sum(a * a, axis=axis)) | |
| elif ord == np.inf: | |
| return np.max(np.abs(a), axis=axis) | |
| else: | |
| raise ValueError(f"Unsupported norm order {ord}. Supported orders are 1, 2, and inf.") | |
| def safe_inv(mat: ndarray, max_retries: int = 4) -> ndarray: | |
| """Compute the inverse of a matrix, no matter it is singular or not. If the matrix is singular, use pseudo-inverse instead. | |
| If both inverse and pseudo-inverse fail, return a matrix filled with NaNs. | |
| ## Parameters | |
| - `mat` (ndarray): shape `(..., M, M)` input square matrix/matrices to invert. | |
| ## Returns | |
| - `inv_mat` (ndarray): shape `(..., M, M)` inverse of the input matrix/matrices. | |
| """ | |
| for i in range(max_retries): | |
| try: | |
| return np.linalg.inv(mat) | |
| except np.linalg.LinAlgError: | |
| eps = 10 ** i * np.finfo(mat.dtype).eps * np.linalg.norm(mat, ord='fro', axis=(-2, -1), keepdims=True) | |
| mat = mat + eps * np.eye(mat.shape[-1]) | |
| try: | |
| return np.linalg.pinv(mat) | |
| except np.linalg.LinAlgError: | |
| warnings.warn("Matrix inversion and pseudo-inversion both failed. Returning NaN matrix.") | |
| return np.full_like(mat, np.nan) | |
| def group(labels: ndarray, data: Optional[np.ndarray] = None) -> List[Tuple[ndarray, ndarray]]: | |
| """ | |
| Split the data into groups based on the provided labels. | |
| ## Parameters | |
| - `labels` `(ndarray)` shape `(N, *label_dims)` array of labels for each data point. Labels can be multi-dimensional. | |
| - `data`: `(ndarray, optional)` shape `(N, *data_dims)` dense tensor. Each one in `N` has `D` features. | |
| If None, return the indices in each group instead. | |
| ## Returns | |
| - `groups` `(List[Tuple[ndarray, ndarray]])`: List of each group, a tuple of `(label, data_in_group)`. | |
| - `label` (ndarray): shape `(*label_dims,)` the label of the group. | |
| - `data_in_group` (ndarray): shape `(length_of_group, *data_dims)` the data points in the group. | |
| If `data` is None, `data_in_group` will be the indices of the data points in the original array. | |
| """ | |
| group_labels, inv, counts = np.unique(labels, return_inverse=True, return_counts=True, axis=0) | |
| if data is None: | |
| data = np.arange(labels.shape[0]) | |
| sections = np.cumsum(counts, axis=0)[:-1] | |
| data_groups = np.split(data[np.argsort(inv)], sections) | |
| return list(zip(group_labels, data_groups)) | |
| def csr_matrix_from_dense_indices(indices: ndarray, n_cols: int) -> 'csr_array': | |
| """Convert a regular indices array to a sparse CSR adjacency matrix format | |
| ## Parameters | |
| - `indices` (ndarray): shape (N, M) dense tensor. Each one in `N` has `M` connections. | |
| - `n_cols` (int): total number of columns in the adjacency matrix | |
| ## Returns | |
| Tensor: shape `(N, n_cols)` sparse CSR adjacency matrix | |
| """ | |
| from scipy.sparse import csr_array | |
| return csr_array(( | |
| np.ones_like(indices, dtype=bool).ravel(), | |
| indices.ravel(), | |
| np.arange(0, indices.size + 1, indices.shape[1]) | |
| ), shape=(indices.shape[0], n_cols)) | |
| def reverse_permutation(perm: ndarray, axis: int = 0) -> ndarray: | |
| """Compute the reverse of a permutation array. | |
| Parameters | |
| ---- | |
| - `perm` (ndarray): shape `(..., N, ...)` permutation array. | |
| - `axis` (int): axis of the permutation array. Other axes are treated as batch dimensions. | |
| Returns | |
| ---- | |
| - `rev_perm` (ndarray): shape `(N,)` reverse permutation array. | |
| Notes | |
| ----- | |
| Equivalent to `np.argsort(perm, axis=axis)`, but more efficient. | |
| """ | |
| axis = axis % perm.ndim | |
| rev_perm = np.empty_like(perm) | |
| indices = np.arange(perm.shape[axis], dtype=perm.dtype)[(None,) * axis + (slice(None),) + (None,) * (perm.ndim - axis - 1)] | |
| np.put_along_axis(rev_perm, perm, indices, axis=axis) | |
| return rev_perm | |
| def vector_outer(x: ndarray, y: Optional[ndarray] = None) -> ndarray: | |
| """ | |
| Compute the outer product of two arrays. | |
| Parameters | |
| ---- | |
| - `x` (ndarray): shape `(..., M)` first array. | |
| - `y` (ndarray, optional): shape `(..., N)` second array. If None, compute the outer product of `x` with itself. | |
| Returns | |
| ---- | |
| - `outer` (ndarray): shape `(..., M, N)` outer product of `x` and `y`. | |
| """ | |
| if y is None: | |
| return x[..., :, None] * x[..., None, :] | |
| return x[..., :, None] * y[..., None, :] |