5cbd3f29e3
Fuzz / Run fuzz harnesses (${{ github.event_name == 'schedule' && 'nightly' || 'smoke' }}) (push) Has been cancelled
Create Releases / call-mac (push) Has been cancelled
Create Releases / call-linux (push) Has been cancelled
Create Releases / call-sdist (push) Has been cancelled
Create Releases / call-win (push) Has been cancelled
Create Releases / call-pyodide (push) Has been cancelled
Windows_No_Exception_CI / build (x64, 3.10) (push) Has been cancelled
Check URLs / build (push) Has been cancelled
Create Releases / Attest CI build artifacts (push) Has been cancelled
Create Releases / Check for Publish release build to pypi (push) Has been cancelled
Create Releases / Check for Publish preview build to test.pypi-weekly (push) Has been cancelled
Create Releases / Publish preview build to test.pypi-weekly (push) Has been cancelled
Create Releases / Check for Publish release build to test.pypi (rc-candidates) (push) Has been cancelled
Create Releases / Publish release build to test.pypi (push) Has been cancelled
Create Releases / Check for Publish preview build to pypi-weekly (push) Has been cancelled
Create Releases / Publish preview build to pypi-weekly (push) Has been cancelled
Create Releases / Publish release build to pypi (push) Has been cancelled
Create Releases / test source distribution (push) Has been cancelled
clang-tidy / clang-tidy (push) Has been cancelled
Lint / Validate SBOM (push) Has been cancelled
Lint / Enforce style (push) Has been cancelled
CI / Test windows-2022, 3.14, External, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test windows-latest, 3.10, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test windows-latest, 3.14, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test windows-latest, 3.14t, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, Internal, debug=1, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, External, debug=0, unity_build=1, onnx_ml=1, autogen=1 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, External, debug=0, unity_build=0, onnx_ml=0, autogen=0 (push) Has been cancelled
CI / Test macos-latest, 3.10, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test macos-latest, 3.14, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test macos-latest, 3.14t, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, External, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.10, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
CI / Test ubuntu-24.04, 3.14t, Internal, debug=0, unity_build=0, onnx_ml=1, autogen=0 (push) Has been cancelled
Pixi CI / Install and lint (ubuntu-24.04-arm) (push) Has been cancelled
Pixi CI / Install and lint (windows-2022) (push) Has been cancelled
Pixi CI / Xcode generator build (push) Has been cancelled
Pixi CI / Install and test (macos-latest, default) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-24.04-arm, default) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-latest, default) (push) Has been cancelled
Pixi CI / Install and test (windows-2022, default) (push) Has been cancelled
Pixi CI / Install and test (macos-latest, oldies) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-24.04-arm, oldies) (push) Has been cancelled
Pixi CI / Install and test (ubuntu-latest, oldies) (push) Has been cancelled
Pixi CI / Install and test (windows-2022, oldies) (push) Has been cancelled
CodeQL / Analyze (actions) (push) Has been cancelled
CodeQL / Analyze (cpp) (push) Has been cancelled
CodeQL / Analyze (python) (push) Has been cancelled
Copilot Setup Steps / copilot-setup-steps (push) Has been cancelled
Generate and publish ONNX docs / build (push) Has been cancelled
Generate and publish ONNX docs / deploy (push) Has been cancelled
Scorecard supply-chain security / Scorecard analysis (push) Has been cancelled
72 lines
2.1 KiB
Python
72 lines
2.1 KiB
Python
# Copyright (c) ONNX Project Contributors
|
|
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
from __future__ import annotations
|
|
|
|
import numpy as np
|
|
|
|
from onnx.reference.op_run import OpRun
|
|
from onnx.reference.ops.op_conv import _conv_implementation
|
|
|
|
|
|
class CausalConvWithState(OpRun):
|
|
def _run(
|
|
self,
|
|
input,
|
|
weight,
|
|
bias=None,
|
|
past_state=None,
|
|
activation=None,
|
|
):
|
|
if activation is None:
|
|
activation = "none"
|
|
if activation not in ("none", "silu", "swish"):
|
|
raise ValueError(
|
|
f"Unsupported activation '{activation}'. "
|
|
"Expected one of: 'none', 'silu', 'swish'."
|
|
)
|
|
|
|
if input.ndim != 3:
|
|
raise ValueError(
|
|
f"input must be rank 3 (batch_size, channels, length), got shape {input.shape}."
|
|
)
|
|
if weight.ndim != 3:
|
|
raise ValueError(
|
|
f"weight must be rank 3 (channels, 1, k), got shape {weight.shape}."
|
|
)
|
|
|
|
batch_size, channels, _ = input.shape
|
|
k = weight.shape[2]
|
|
|
|
# Step 1: build the left-padded input (B, C, L + k - 1).
|
|
if past_state is None:
|
|
pad = np.zeros((batch_size, channels, k - 1), dtype=input.dtype)
|
|
else:
|
|
pad = past_state
|
|
padded = np.concatenate([pad, input], axis=2)
|
|
|
|
# Step 2: depthwise Conv1d (group = channels, valid padding).
|
|
conv_out = _conv_implementation(
|
|
padded,
|
|
weight,
|
|
bias,
|
|
"NOTSET",
|
|
[1],
|
|
channels,
|
|
[k],
|
|
[0, 0],
|
|
[1],
|
|
).astype(input.dtype)
|
|
|
|
# Step 3: optional fused SiLU/Swish activation.
|
|
if activation in ("silu", "swish"):
|
|
sigmoid = 1.0 / (1.0 + np.exp(-conv_out.astype(np.float32)))
|
|
output = (conv_out.astype(np.float32) * sigmoid).astype(input.dtype)
|
|
else:
|
|
output = conv_out
|
|
|
|
# Step 4: present_state = last (k - 1) positions of the padded input.
|
|
present_state = padded[:, :, padded.shape[2] - (k - 1) :]
|
|
|
|
return (output, present_state)
|