diff --git a/tinygrad/nn/onnx.py b/tinygrad/nn/onnx.py index 4083e3be75..7723e74b3e 100644 --- a/tinygrad/nn/onnx.py +++ b/tinygrad/nn/onnx.py @@ -556,8 +556,8 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT return tuple(t if cond else e for t,e in zip(then_out.values(), else_out.values())) def Identity(x:Tensor): return x - def Constant(sparse_value:Tensor|None=None, value:Tensor|None=None, value_float:float|None=None, value_floats:list[float]|None=None, - value_int:int|None=None, value_ints:list[int]|None=None, value_string:str|None=None, value_strings:list[str]|None=None): + def Constant(sparse_value:Tensor|None=None, value:Tensor|None=None, value_float:float|None=None, value_floats:tuple[float, ...]|None=None, + value_int:int|None=None, value_ints:tuple[int, ...]|None=None, value_string:str|None=None, value_strings:tuple[str, ...]|None=None): if value is not None: return value if value_float is not None: return Tensor(value_float, dtype=dtypes.float32) if value_floats is not None: return Tensor(list(value_floats), dtype=dtypes.float32) @@ -594,7 +594,7 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT # ***** Unary Ops (math) ***** def Not(x:Tensor): return x.logical_not() - def Clip(x: Tensor, min:Tensor|None=None, max:Tensor|None=None): return x if min is None and max is None else x.clip(min, max) # noqa: A002 # pylint: disable=redefined-builtin + def Clip(x: Tensor, min:Tensor|float|None=None, max:Tensor|float|None=None): return x if min is None and max is None else x.clip(min, max) # noqa: A002 # pylint: disable=redefined-builtin def IsInf(x:Tensor, detect_negative:int=1, detect_positive:int=1): return x.isinf(bool(detect_positive), bool(detect_negative)) # ***** Unary Ops (activation) ***** @@ -643,26 +643,26 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT def Min(*data_0:Tensor): return functools.reduce(Tensor.minimum, data_0) def Sum(*data_0:Tensor): return functools.reduce(Tensor.add, data_0) def Mean(*data_0:Tensor): return Sum(*data_0) / len(data_0) - def ReduceMax(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceMax(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return data.max(_axes(axes, noop_with_empty_axes), keepdim=keepdims) - def ReduceMin(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceMin(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return data.min(_axes(axes, noop_with_empty_axes), keepdim=keepdims) - def ReduceSum(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceSum(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return data.sum(_axes(axes, noop_with_empty_axes), keepdim=keepdims) - def ReduceMean(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceMean(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return data.mean(_axes(axes, noop_with_empty_axes), keepdim=keepdims) - def ReduceSumSquare(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceSumSquare(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return ReduceSum(data.square(), axes, keepdims, noop_with_empty_axes) - def ReduceProd(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceProd(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return data.prod(_axes(axes, noop_with_empty_axes), keepdim=keepdims) - def ReduceL1(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceL1(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return ReduceSum(data.abs(), axes, keepdims, noop_with_empty_axes) - def ReduceL2(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceL2(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): dtype = dtypes.float if data.dtype in (dtypes.float16, dtypes.bfloat16) else data.dtype return ReduceSum(data.cast(dtype).square(), axes, keepdims, noop_with_empty_axes).sqrt().cast(data.dtype) - def ReduceLogSum(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceLogSum(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return ReduceSum(data, axes, keepdims, noop_with_empty_axes).log() - def ReduceLogSumExp(data:Tensor, axes:list[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): + def ReduceLogSumExp(data:Tensor, axes:Sequence[int]|None=None, keepdims:int=1, noop_with_empty_axes:int=0): return ReduceSum(data.exp(), axes, keepdims, noop_with_empty_axes).log() def ArgMax(x:Tensor, axis:int=0, keepdims:int=1, select_last_index:int=0): if select_last_index: return ((int(x.shape[axis])-1) - x.flip(axis).argmax(axis, keepdim=keepdims)).cast(dtypes.int64) @@ -671,32 +671,32 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT return ArgMax(-x, axis=axis, keepdims=keepdims, select_last_index=select_last_index) # ***** Movement Ops ***** - def Reshape(data:Tensor, shape:list[int], allowzero:int=0): + def Reshape(data:Tensor, shape:Sequence[int], allowzero:int=0): return data.reshape([x if x != 0 else (0 if allowzero else data.shape[i]) for i,x in enumerate(shape)]) def Flatten(x:Tensor, axis:int=1): return x.reshape(prod(x.shape[0:axis]), -1) def Expand(x:Tensor, shape:list[int]): return x.expand(_broadcast_shape(x.shape, tuple(shape))) def Shrink(x:Tensor, bias:float=0.0, lambd:float=0.5): return (x < -lambd)*(x+bias) + (x > lambd)*(x-bias) - def Transpose(x:Tensor, perm:list[int]|None=None): return x.permute(order=perm or list(range(x.ndim)[::-1])) + def Transpose(x:Tensor, perm:tuple[int, ...]|None=None): return x.permute(order=perm or list(range(x.ndim)[::-1])) - def Squeeze(data:Tensor, axes:list[int]|None=None): + def Squeeze(data:Tensor, axes:Sequence[int]|None=None): return data.squeeze() if axes is None else functools.reduce(lambda d, dim: d.squeeze(dim), sorted(axes, reverse=True), data) - def Unsqueeze(data:Tensor, axes:list[int]): return functools.reduce(lambda d, dim: d.unsqueeze(dim), sorted(axes), data) + def Unsqueeze(data:Tensor, axes:Sequence[int]): return functools.reduce(lambda d, dim: d.unsqueeze(dim), sorted(axes), data) def Tile(x:Tensor, repeats:list[int]): return x.repeat(repeats) def Concat(*xs:Tensor, axis:int): return Tensor.cat(*xs, dim=axis) - def Slice(data:Tensor, starts:list[int], ends:list[int], axes:list[int]|None=None, steps:list[int]|None=None): + def Slice(data:Tensor, starts:Sequence[int], ends:Sequence[int], axes:Sequence[int]|None=None, steps:list[int]|None=None): axes = axes or list(range(data.ndim)) steps = steps or [1] * data.ndim slices = [slice(None)] * data.ndim for i, axis in enumerate(axes): slices[axis] = slice(starts[i], ends[i], steps[i]) return data[tuple(slices)] - def Split(data:Tensor, split:list[int]|None=None, num_outputs:int=0, axis:int=0): + def Split(data:Tensor, split:Sequence[int]|None=None, num_outputs:int=0, axis:int=0): sz = int(data.shape[axis]) if split is None: split = [sz // num_outputs + (1 if i < sz % num_outputs else 0) for i in range(num_outputs)] return data.split(split, axis) - def Pad(x:Tensor, pads:list[int], constant_value:ConstType|None=None, axes:list[int]|None=None, + def Pad(x:Tensor, pads:Sequence[int], constant_value:ConstType|None=None, axes:list[int]|None=None, mode:Literal["constant", "reflect", "edge", "wrap"]="constant", value=0): value = _resolve_const(value if constant_value is None else constant_value) axes = axes or list(range(x.ndim)) @@ -704,7 +704,7 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT for i,axis in enumerate(axes): real_pads[axis%x.ndim], real_pads[axis%x.ndim+x.ndim] = pads[i], pads[i+len(axes)] return x.pad(padding=_onnx_pads_to_tiny_pads(real_pads), mode={"edge":"replicate", "wrap":"circular"}.get(mode, mode), value=value) - def CenterCropPad(t:Tensor, shape:list[int], axes:list[int]|None=None): + def CenterCropPad(t:Tensor, shape:list[int], axes:tuple[int, ...]|None=None): shrink_arg:list[None|tuple[sint,sint]] = [None] * t.ndim pad_arg:list[None|tuple[sint,sint]] = [None] * t.ndim for s, x in zip(shape, axes or range(t.ndim)): @@ -714,26 +714,26 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT return t.shrink(tuple(shrink_arg)).pad(tuple(pad_arg)) # ***** Processing Ops ***** - def AveragePool(X: Tensor, kernel_shape:list[int], auto_pad:AUTO_PAD_OPTIONS="NOTSET", ceil_mode:int=0, count_include_pad:int=0, - dilations:list[int]|int=1, pads:list[int]|int=0, strides:list[int]|int=1): + def AveragePool(X: Tensor, kernel_shape:tuple[int, ...], auto_pad:AUTO_PAD_OPTIONS="NOTSET", ceil_mode:int=0, count_include_pad:int=0, + dilations:tuple[int, ...]|int=1, pads:tuple[int, ...]|int=0, strides:tuple[int, ...]|int=1): pool_pads = _resolve_pool_pads(X, pads, kernel_shape, dilations, strides, auto_pad) return X.avg_pool2d(tuple(kernel_shape), strides, dilations, pool_pads, ceil_mode=ceil_mode, count_include_pad=count_include_pad) - def MaxPool(X: Tensor, kernel_shape:list[int], auto_pad:AUTO_PAD_OPTIONS="NOTSET", ceil_mode:int=0, dilations:list[int]|int=1, pads:list[int]|int=0, - storage_order:int=0, strides:list[int]|int=1): + def MaxPool(X: Tensor, kernel_shape:tuple[int, ...], auto_pad:AUTO_PAD_OPTIONS="NOTSET", ceil_mode:int=0, dilations:tuple[int, ...]|int=1, + pads:tuple[int, ...]|int=0, storage_order:int=0, strides:tuple[int, ...]|int=1): pool_pads = _resolve_pool_pads(X, pads, kernel_shape, dilations, strides, auto_pad) out = X.max_pool2d(tuple(kernel_shape), strides, dilations, pool_pads, ceil_mode=ceil_mode, return_indices=True) ret, idx = cast(tuple[Tensor, Tensor], out) return ret, idx.transpose(-2, -1).cast(dtypes.int64) if storage_order else idx.cast(dtypes.int64) - def Conv(X: Tensor, W: Tensor, B:Tensor|None=None, auto_pad:AUTO_PAD_OPTIONS="NOTSET", dilations:list[int]|int=1, group:int=1, - kernel_shape:list[int]|None=None, pads:list[int]|int=0, strides:list[int]|int=1): + def Conv(X: Tensor, W: Tensor, B:Tensor|None=None, auto_pad:AUTO_PAD_OPTIONS="NOTSET", dilations:tuple[int, ...]|int=1, group:int=1, + kernel_shape:tuple[int, ...]|None=None, pads:tuple[int, ...]|int=0, strides:tuple[int, ...]|int=1): return X.conv2d(W, B, stride=strides, groups=group, dilation=dilations, padding=_resolve_pool_pads(X, pads, kernel_shape or W.shape[2:], dilations, strides, auto_pad)) - def ConvTranspose(X: Tensor, W: Tensor, B:Tensor|None=None, auto_pad:AUTO_PAD_OPTIONS="NOTSET", dilations:list[int]|int=1, group:int=1, - kernel_shape:list[int]|None=None, pads:list[int]|None=None, output_shape:list[int]|None=None, output_padding:list[int]|int=0, - strides:list[int]|int=1): + def ConvTranspose(X: Tensor, W: Tensor, B:Tensor|None=None, auto_pad:AUTO_PAD_OPTIONS="NOTSET", dilations:tuple[int, ...]|int=1, group:int=1, + kernel_shape:tuple[int, ...]|None=None, pads:Sequence[int]|None=None, output_shape:Sequence[int]|None=None, + output_padding:tuple[int, ...]|int=0, strides:tuple[int, ...]|int=1): input_shape_, kernel_shape_ = X.shape[2:], (kernel_shape or W.shape[2:]) strides_, dilations_, output_padding_ = (make_tuple(x, len(input_shape_)) for x in (strides, dilations, output_padding)) if output_shape is not None: # we pad according to output_shape @@ -747,8 +747,8 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT pads = _onnx_pads_to_tiny_pads(pads) return X.conv_transpose2d(W, B, group, strides_, dilations_, pads, output_padding_) - def MaxUnpool(xT: Tensor, xI: Tensor, outshape: list[int]|None=None, kernel_shape:list[int]|None=None, pads:list[int]|int=0, - strides:list[int]|int=1): + def MaxUnpool(xT: Tensor, xI: Tensor, outshape: list[int]|None=None, kernel_shape:Sequence[int]|None=None, pads:tuple[int, ...]|int=0, + strides:tuple[int, ...]|int=1): if kernel_shape is None: kernel_shape = [] pads_: int | tuple[int, ...] = pads if isinstance(pads, int) else _onnx_pads_to_tiny_pads(pads) return Tensor.max_unpool2d(xT, xI, tuple(kernel_shape), strides, 1, pads_, outshape if outshape is None else tuple(outshape)) @@ -761,7 +761,7 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT if C is not None: ret = ret + beta * (C if broadcast == 0 else C.reshape([-1 if i < len(C.shape) else 1 for i in range(ret.ndim)][::-1])) return ret - def Einsum(*Inputs:list[Tensor], equation:str): return Tensor.einsum(equation, *Inputs) + def Einsum(*Inputs:Tensor, equation:str): return Tensor.einsum(equation, *Inputs) def CumSum(X:Tensor, axis:int|list[int], exclusive:int=0, reverse:int=0): axis = X._resolve_dim(_resolve_const(axis)) @@ -774,8 +774,8 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT k_ = _resolve_const(k) return x.triu(k_) if upper else x.tril(k_) - def Resize(X:Tensor, roi:list[float]|None=None, scales:list[float]|None=None, sizes:list[int]|None=None, antialias:int=0, - axes:list[int]|None=None, coordinate_transformation_mode:str='half_pixel', cubic_coeff_a:float=-0.75, exclude_outside:int=0, + def Resize(X:Tensor, roi:list[float]|None=None, scales:Sequence[float]|None=None, sizes:list[int]|None=None, antialias:int=0, + axes:Sequence[int]|None=None, coordinate_transformation_mode:str='half_pixel', cubic_coeff_a:float=-0.75, exclude_outside:int=0, extrapolation_value:float=0.0, keep_aspect_ratio_policy:str='stretch', mode:str='nearest', nearest_mode:str='round_prefer_floor'): def _apply_transformation(input_sz, output_sz, scale_dim, mode): index = Tensor.arange(output_sz) @@ -876,7 +876,7 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT gathered_values = [X.gather(i, idx) for idx in expanded_indices] X = sum(v * c for v, c in zip(gathered_values, expanded_coeffs)) return X.permute(*argsort(perm)) if perm else X - def Upsample(X, scales, mode): return Resize(X=X, scales=scales, mode=mode) # deprecated + def Upsample(X:Tensor, scales:Sequence[float], mode:str): return Resize(X=X, scales=scales, mode=mode) # deprecated def TopK(X:Tensor, K:int|list[int], axis:int=-1, largest:int=1, sorted:int=1): # noqa: A002 # pylint: disable=redefined-builtin val, idx = X.topk(_resolve_const(K), axis, bool(largest), bool(sorted)) @@ -937,7 +937,7 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT if segment_embedding is not None: embedding_sum = embedding_sum + embedding(segment_ids, segment_embedding.shape[0], segment_embedding) out = embedding_sum.layernorm(eps=epsilon) * gamma + beta return out, None, embedding_sum - def MeanVarianceNormalization(x:Tensor, axis:list[int]|None=None): + def MeanVarianceNormalization(x:Tensor, axis:Sequence[int]|None=None): if axis is None: axis = [0,2,3] return (x - x.mean(axis, keepdim=True)) / (x.std(axis, keepdim=True, correction=0) + 1e-9) @@ -1001,7 +1001,7 @@ def get_onnx_ops() -> dict[str, types.FunctionType|dict[OpSetId, types.FunctionT def attention_contrib(x:Tensor, weights:Tensor, bias:Tensor|None=None, mask_index:Tensor|None=None, past:Tensor|None=None, attention_bias:Tensor|None=None, past_sequence_length:Tensor|None=None, do_rotary:int=0, mask_filter_value:float=-10000.0, - num_heads:int|None=None, past_present_share_buffer:int|None=None, qkv_hidden_sizes:list[int]|None=None, + num_heads:int|None=None, past_present_share_buffer:int|None=None, qkv_hidden_sizes:Sequence[int]|None=None, rotary_embedding_dim:int|None=None, scale:float|None=None, unidirectional:int=0): assert not do_rotary and not attention_bias, "TODO" if qkv_hidden_sizes is None: qkv_hidden_sizes = [int(weights.shape[1] // 3)] * 3