pruned-ctc 0.1.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,59 @@
1
+ # Contributing
2
+
3
+ ## Development
4
+
5
+ Use Python 3.10+. Linting, formatting, and package builds do not require
6
+ PyTorch or k2. Runtime installation is described in [README.md](README.md).
7
+
8
+ ```bash
9
+ python -m pip install --upgrade ruff
10
+ python -m ruff check .
11
+ python -m ruff format --check .
12
+ ```
13
+
14
+ CI checks linting, formatting, and package builds.
15
+
16
+ Keep the public API focused on `pruned_ctc_loss`.
17
+ Contributions use the [MIT license](LICENSE).
18
+
19
+ ## Building a package
20
+
21
+ The build does not require PyTorch or k2:
22
+
23
+ ```bash
24
+ python -m pip install --upgrade build twine
25
+ python -m build
26
+ python -m twine check --strict dist/*
27
+ ```
28
+
29
+ This produces a wheel and source distribution in `dist/`. CI also builds,
30
+ checks, and uploads these files as workflow artifacts.
31
+
32
+ ## Publishing a release
33
+
34
+ Publishing is manual. The `Publish` workflow runs only through
35
+ `workflow_dispatch`; pushes and GitHub releases do not publish packages.
36
+
37
+ Before the first upload, the repository owner must configure
38
+ [Trusted Publishing](https://docs.pypi.org/trusted-publishers/) independently on
39
+ PyPI and TestPyPI and create the corresponding GitHub environments:
40
+
41
+ | Setting | TestPyPI | PyPI |
42
+ | --- | --- | --- |
43
+ | Distribution | `pruned-ctc` | `pruned-ctc` |
44
+ | Repository owner | `yfyeung` | `yfyeung` |
45
+ | Repository | `PrunedCTC` | `PrunedCTC` |
46
+ | Workflow filename | `publish.yml` | `publish.yml` |
47
+ | GitHub environment | `testpypi` | `pypi` |
48
+
49
+ A new project can use a pending Trusted Publisher. Configure environment
50
+ protection rules as appropriate for the repository. Authentication uses a
51
+ short-lived GitHub OIDC token; no PyPI password or API token is stored in the
52
+ workflow.
53
+
54
+ For each release, update `__version__` in `pruned_ctc.py`, run the checks above,
55
+ and push the reviewed revision. In GitHub Actions, select **Publish → Run
56
+ workflow** and choose that revision. Start with the default `testpypi` target;
57
+ after checking its uploaded distributions, run the workflow for `pypi` using
58
+ the same revision. Each run builds and checks its artifacts before the
59
+ publishing job enters the selected environment.
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 Shanghai Jiao Tong University
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1 @@
1
+ include LICENSE README.md CONTRIBUTING.md
@@ -0,0 +1,150 @@
1
+ Metadata-Version: 2.4
2
+ Name: pruned-ctc
3
+ Version: 0.1.0
4
+ Summary: Memory-efficient CTC loss with exact vocabulary reduction and pruned alignments
5
+ Author: Yifan Yang
6
+ License-Expression: MIT
7
+ Project-URL: Homepage, https://github.com/yfyeung/PrunedCTC
8
+ Project-URL: Repository, https://github.com/yfyeung/PrunedCTC
9
+ Project-URL: Issues, https://github.com/yfyeung/PrunedCTC/issues
10
+ Keywords: ctc,speech-recognition,pytorch,sequence-modeling
11
+ Classifier: Development Status :: 4 - Beta
12
+ Classifier: Intended Audience :: Science/Research
13
+ Classifier: Programming Language :: Python :: 3
14
+ Classifier: Programming Language :: Python :: 3 :: Only
15
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
16
+ Requires-Python: >=3.10
17
+ Description-Content-Type: text/markdown
18
+ License-File: LICENSE
19
+ Requires-Dist: torch>=2.4
20
+ Provides-Extra: dev
21
+ Requires-Dist: build>=1.2.2; extra == "dev"
22
+ Requires-Dist: ruff>=0.9; extra == "dev"
23
+ Requires-Dist: twine>=6; extra == "dev"
24
+ Dynamic: license-file
25
+
26
+ # Pruned CTC
27
+
28
+ Memory-efficient CTC training for large vocabularies, implemented in PyTorch and
29
+ [k2](https://github.com/k2-fsa/k2).
30
+
31
+ Pruned CTC computes the CTC loss directly from encoder states and a linear
32
+ projection, avoiding a full `(batch, time, vocabulary)` logit tensor. It combines:
33
+
34
+ - **Exact vocabulary reduction:** the CTC graph uses only the target tokens and
35
+ blank. Probabilities are still normalized over the **full vocabulary**, and
36
+ gradients include every vocabulary class.
37
+ - **Chunked projection and backward:** vocabulary chunks are recomputed during
38
+ backward to reduce the memory used by the projection and loss activations.
39
+ - **Alignment pruning:** k2 prunes the alignment lattice using a configurable
40
+ log-score beam. A finite beam makes the alignment sum approximate; vocabulary reduction
41
+ itself does not introduce this approximation.
42
+
43
+ The projection weights, their gradients, and optimizer states remain dense.
44
+ Memory savings depend on the vocabulary, batch, sequence lengths, chunk size,
45
+ and alignment beam.
46
+
47
+ ## Installation
48
+
49
+ Requires **Python 3.10+**, **PyTorch 2.4+**, and a compatible **k2** build.
50
+ First install PyTorch and k2 for your Python version, device, and CUDA version.
51
+ Follow the
52
+ [k2 installation guide](https://k2-fsa.github.io/k2/installation/index.html);
53
+ k2 wheels are tied to specific PyTorch builds. A generic `pip install k2` can
54
+ select an older PyTorch dependency and replace an existing installation.
55
+
56
+ Install Pruned CTC from PyPI:
57
+
58
+ ```bash
59
+ python -m pip install pruned-ctc
60
+ ```
61
+
62
+ Or install from this repository:
63
+
64
+ ```bash
65
+ git clone https://github.com/yfyeung/PrunedCTC.git
66
+ cd PrunedCTC
67
+ python -m pip install .
68
+ ```
69
+
70
+ ## Quick start
71
+
72
+ ```python
73
+ import k2
74
+ import torch
75
+
76
+ from pruned_ctc import pruned_ctc_loss
77
+
78
+ device = torch.device("cpu") # Use "cuda" with a compatible CUDA-enabled k2.
79
+
80
+ encoder_out = torch.randn(2, 12, 8, device=device, requires_grad=True)
81
+ head = torch.nn.Linear(8, 6, device=device)
82
+ targets = k2.RaggedTensor([[1, 2, 3], [2, 2]]).to(device)
83
+ lengths = torch.tensor([12, 9], dtype=torch.int64, device=device)
84
+
85
+ with torch.autocast(device.type, enabled=False):
86
+ loss = pruned_ctc_loss(
87
+ encoder_out=encoder_out,
88
+ weight=head.weight,
89
+ bias=head.bias,
90
+ y=targets,
91
+ encoder_out_lens=lengths,
92
+ chunk_size=4,
93
+ output_beam=100.0,
94
+ blank_id=0,
95
+ )
96
+
97
+ loss.backward()
98
+ print(f"Summed CTC loss: {loss.item():.4f}")
99
+ ```
100
+
101
+ ## API
102
+
103
+ ```python
104
+ pruned_ctc_loss(
105
+ encoder_out,
106
+ weight,
107
+ bias,
108
+ y,
109
+ encoder_out_lens,
110
+ chunk_size=4096,
111
+ output_beam=100.0,
112
+ max_states=100_000_000,
113
+ blank_id=0,
114
+ )
115
+ ```
116
+
117
+ | Argument | Expected value |
118
+ | --- | --- |
119
+ | `encoder_out` | Encoder states of shape `(N, T, D)`. |
120
+ | `weight` | Linear projection weights of shape `(V, D)`. |
121
+ | `bias` | Projection bias of shape `(V,)`, or `None`. |
122
+ | `y` | Two-axis `k2.RaggedTensor`: one sequence of target IDs per utterance, with no blank tokens. |
123
+ | `encoder_out_lens` | Int32/int64 frame counts of shape `(N,)`, each between `1` and `T`. |
124
+ | `chunk_size` | Number of vocabulary columns processed per chunk. Smaller chunks reduce temporary projection memory and can cost speed. |
125
+ | `output_beam` | Positive, finite log-score beam used for alignment pruning. Larger beams retain more alignments and use more memory. |
126
+ | `max_states` | Lattice state limit. k2 may narrow the effective beam to stay within this limit. |
127
+ | `blank_id` | Blank token ID in the original vocabulary; it is remapped internally to column zero for k2. |
128
+
129
+ The return value is a scalar `float32` **sum** over utterances. Nonfinite
130
+ per-utterance losses are excluded with a warning, including losses from
131
+ impossible target alignments. With finite inputs, impossible alignments
132
+ contribute zero loss and zero gradients. Apply any desired loss normalization
133
+ explicitly.
134
+
135
+ ### Precision and device requirements
136
+
137
+ - `encoder_out`, `weight`, and `bias` may each use `float32` or `bfloat16`.
138
+ Their values must be finite. Gradients retain the corresponding input dtypes.
139
+ `float16` is unsupported.
140
+ - Disable autocast around the loss call, even when the encoder runs under mixed
141
+ precision. The loss manages its own numerical precision.
142
+ - Put the encoder states, projection parameters, and targets on the same CPU or
143
+ CUDA device. Frame counts may also stay on the CPU.
144
+ - Target IDs must be valid vocabulary indices and must exclude `blank_id`.
145
+ - Only first-order gradients are supported; double backward is unsupported.
146
+
147
+ ## License
148
+
149
+ [MIT](https://github.com/yfyeung/PrunedCTC/blob/master/LICENSE). Copyright 2026 Shanghai Jiao Tong University.
150
+ Author: Yifan Yang.
@@ -0,0 +1,125 @@
1
+ # Pruned CTC
2
+
3
+ Memory-efficient CTC training for large vocabularies, implemented in PyTorch and
4
+ [k2](https://github.com/k2-fsa/k2).
5
+
6
+ Pruned CTC computes the CTC loss directly from encoder states and a linear
7
+ projection, avoiding a full `(batch, time, vocabulary)` logit tensor. It combines:
8
+
9
+ - **Exact vocabulary reduction:** the CTC graph uses only the target tokens and
10
+ blank. Probabilities are still normalized over the **full vocabulary**, and
11
+ gradients include every vocabulary class.
12
+ - **Chunked projection and backward:** vocabulary chunks are recomputed during
13
+ backward to reduce the memory used by the projection and loss activations.
14
+ - **Alignment pruning:** k2 prunes the alignment lattice using a configurable
15
+ log-score beam. A finite beam makes the alignment sum approximate; vocabulary reduction
16
+ itself does not introduce this approximation.
17
+
18
+ The projection weights, their gradients, and optimizer states remain dense.
19
+ Memory savings depend on the vocabulary, batch, sequence lengths, chunk size,
20
+ and alignment beam.
21
+
22
+ ## Installation
23
+
24
+ Requires **Python 3.10+**, **PyTorch 2.4+**, and a compatible **k2** build.
25
+ First install PyTorch and k2 for your Python version, device, and CUDA version.
26
+ Follow the
27
+ [k2 installation guide](https://k2-fsa.github.io/k2/installation/index.html);
28
+ k2 wheels are tied to specific PyTorch builds. A generic `pip install k2` can
29
+ select an older PyTorch dependency and replace an existing installation.
30
+
31
+ Install Pruned CTC from PyPI:
32
+
33
+ ```bash
34
+ python -m pip install pruned-ctc
35
+ ```
36
+
37
+ Or install from this repository:
38
+
39
+ ```bash
40
+ git clone https://github.com/yfyeung/PrunedCTC.git
41
+ cd PrunedCTC
42
+ python -m pip install .
43
+ ```
44
+
45
+ ## Quick start
46
+
47
+ ```python
48
+ import k2
49
+ import torch
50
+
51
+ from pruned_ctc import pruned_ctc_loss
52
+
53
+ device = torch.device("cpu") # Use "cuda" with a compatible CUDA-enabled k2.
54
+
55
+ encoder_out = torch.randn(2, 12, 8, device=device, requires_grad=True)
56
+ head = torch.nn.Linear(8, 6, device=device)
57
+ targets = k2.RaggedTensor([[1, 2, 3], [2, 2]]).to(device)
58
+ lengths = torch.tensor([12, 9], dtype=torch.int64, device=device)
59
+
60
+ with torch.autocast(device.type, enabled=False):
61
+ loss = pruned_ctc_loss(
62
+ encoder_out=encoder_out,
63
+ weight=head.weight,
64
+ bias=head.bias,
65
+ y=targets,
66
+ encoder_out_lens=lengths,
67
+ chunk_size=4,
68
+ output_beam=100.0,
69
+ blank_id=0,
70
+ )
71
+
72
+ loss.backward()
73
+ print(f"Summed CTC loss: {loss.item():.4f}")
74
+ ```
75
+
76
+ ## API
77
+
78
+ ```python
79
+ pruned_ctc_loss(
80
+ encoder_out,
81
+ weight,
82
+ bias,
83
+ y,
84
+ encoder_out_lens,
85
+ chunk_size=4096,
86
+ output_beam=100.0,
87
+ max_states=100_000_000,
88
+ blank_id=0,
89
+ )
90
+ ```
91
+
92
+ | Argument | Expected value |
93
+ | --- | --- |
94
+ | `encoder_out` | Encoder states of shape `(N, T, D)`. |
95
+ | `weight` | Linear projection weights of shape `(V, D)`. |
96
+ | `bias` | Projection bias of shape `(V,)`, or `None`. |
97
+ | `y` | Two-axis `k2.RaggedTensor`: one sequence of target IDs per utterance, with no blank tokens. |
98
+ | `encoder_out_lens` | Int32/int64 frame counts of shape `(N,)`, each between `1` and `T`. |
99
+ | `chunk_size` | Number of vocabulary columns processed per chunk. Smaller chunks reduce temporary projection memory and can cost speed. |
100
+ | `output_beam` | Positive, finite log-score beam used for alignment pruning. Larger beams retain more alignments and use more memory. |
101
+ | `max_states` | Lattice state limit. k2 may narrow the effective beam to stay within this limit. |
102
+ | `blank_id` | Blank token ID in the original vocabulary; it is remapped internally to column zero for k2. |
103
+
104
+ The return value is a scalar `float32` **sum** over utterances. Nonfinite
105
+ per-utterance losses are excluded with a warning, including losses from
106
+ impossible target alignments. With finite inputs, impossible alignments
107
+ contribute zero loss and zero gradients. Apply any desired loss normalization
108
+ explicitly.
109
+
110
+ ### Precision and device requirements
111
+
112
+ - `encoder_out`, `weight`, and `bias` may each use `float32` or `bfloat16`.
113
+ Their values must be finite. Gradients retain the corresponding input dtypes.
114
+ `float16` is unsupported.
115
+ - Disable autocast around the loss call, even when the encoder runs under mixed
116
+ precision. The loss manages its own numerical precision.
117
+ - Put the encoder states, projection parameters, and targets on the same CPU or
118
+ CUDA device. Frame counts may also stay on the CPU.
119
+ - Target IDs must be valid vocabulary indices and must exclude `blank_id`.
120
+ - Only first-order gradients are supported; double backward is unsupported.
121
+
122
+ ## License
123
+
124
+ [MIT](https://github.com/yfyeung/PrunedCTC/blob/master/LICENSE). Copyright 2026 Shanghai Jiao Tong University.
125
+ Author: Yifan Yang.
@@ -0,0 +1,150 @@
1
+ Metadata-Version: 2.4
2
+ Name: pruned-ctc
3
+ Version: 0.1.0
4
+ Summary: Memory-efficient CTC loss with exact vocabulary reduction and pruned alignments
5
+ Author: Yifan Yang
6
+ License-Expression: MIT
7
+ Project-URL: Homepage, https://github.com/yfyeung/PrunedCTC
8
+ Project-URL: Repository, https://github.com/yfyeung/PrunedCTC
9
+ Project-URL: Issues, https://github.com/yfyeung/PrunedCTC/issues
10
+ Keywords: ctc,speech-recognition,pytorch,sequence-modeling
11
+ Classifier: Development Status :: 4 - Beta
12
+ Classifier: Intended Audience :: Science/Research
13
+ Classifier: Programming Language :: Python :: 3
14
+ Classifier: Programming Language :: Python :: 3 :: Only
15
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
16
+ Requires-Python: >=3.10
17
+ Description-Content-Type: text/markdown
18
+ License-File: LICENSE
19
+ Requires-Dist: torch>=2.4
20
+ Provides-Extra: dev
21
+ Requires-Dist: build>=1.2.2; extra == "dev"
22
+ Requires-Dist: ruff>=0.9; extra == "dev"
23
+ Requires-Dist: twine>=6; extra == "dev"
24
+ Dynamic: license-file
25
+
26
+ # Pruned CTC
27
+
28
+ Memory-efficient CTC training for large vocabularies, implemented in PyTorch and
29
+ [k2](https://github.com/k2-fsa/k2).
30
+
31
+ Pruned CTC computes the CTC loss directly from encoder states and a linear
32
+ projection, avoiding a full `(batch, time, vocabulary)` logit tensor. It combines:
33
+
34
+ - **Exact vocabulary reduction:** the CTC graph uses only the target tokens and
35
+ blank. Probabilities are still normalized over the **full vocabulary**, and
36
+ gradients include every vocabulary class.
37
+ - **Chunked projection and backward:** vocabulary chunks are recomputed during
38
+ backward to reduce the memory used by the projection and loss activations.
39
+ - **Alignment pruning:** k2 prunes the alignment lattice using a configurable
40
+ log-score beam. A finite beam makes the alignment sum approximate; vocabulary reduction
41
+ itself does not introduce this approximation.
42
+
43
+ The projection weights, their gradients, and optimizer states remain dense.
44
+ Memory savings depend on the vocabulary, batch, sequence lengths, chunk size,
45
+ and alignment beam.
46
+
47
+ ## Installation
48
+
49
+ Requires **Python 3.10+**, **PyTorch 2.4+**, and a compatible **k2** build.
50
+ First install PyTorch and k2 for your Python version, device, and CUDA version.
51
+ Follow the
52
+ [k2 installation guide](https://k2-fsa.github.io/k2/installation/index.html);
53
+ k2 wheels are tied to specific PyTorch builds. A generic `pip install k2` can
54
+ select an older PyTorch dependency and replace an existing installation.
55
+
56
+ Install Pruned CTC from PyPI:
57
+
58
+ ```bash
59
+ python -m pip install pruned-ctc
60
+ ```
61
+
62
+ Or install from this repository:
63
+
64
+ ```bash
65
+ git clone https://github.com/yfyeung/PrunedCTC.git
66
+ cd PrunedCTC
67
+ python -m pip install .
68
+ ```
69
+
70
+ ## Quick start
71
+
72
+ ```python
73
+ import k2
74
+ import torch
75
+
76
+ from pruned_ctc import pruned_ctc_loss
77
+
78
+ device = torch.device("cpu") # Use "cuda" with a compatible CUDA-enabled k2.
79
+
80
+ encoder_out = torch.randn(2, 12, 8, device=device, requires_grad=True)
81
+ head = torch.nn.Linear(8, 6, device=device)
82
+ targets = k2.RaggedTensor([[1, 2, 3], [2, 2]]).to(device)
83
+ lengths = torch.tensor([12, 9], dtype=torch.int64, device=device)
84
+
85
+ with torch.autocast(device.type, enabled=False):
86
+ loss = pruned_ctc_loss(
87
+ encoder_out=encoder_out,
88
+ weight=head.weight,
89
+ bias=head.bias,
90
+ y=targets,
91
+ encoder_out_lens=lengths,
92
+ chunk_size=4,
93
+ output_beam=100.0,
94
+ blank_id=0,
95
+ )
96
+
97
+ loss.backward()
98
+ print(f"Summed CTC loss: {loss.item():.4f}")
99
+ ```
100
+
101
+ ## API
102
+
103
+ ```python
104
+ pruned_ctc_loss(
105
+ encoder_out,
106
+ weight,
107
+ bias,
108
+ y,
109
+ encoder_out_lens,
110
+ chunk_size=4096,
111
+ output_beam=100.0,
112
+ max_states=100_000_000,
113
+ blank_id=0,
114
+ )
115
+ ```
116
+
117
+ | Argument | Expected value |
118
+ | --- | --- |
119
+ | `encoder_out` | Encoder states of shape `(N, T, D)`. |
120
+ | `weight` | Linear projection weights of shape `(V, D)`. |
121
+ | `bias` | Projection bias of shape `(V,)`, or `None`. |
122
+ | `y` | Two-axis `k2.RaggedTensor`: one sequence of target IDs per utterance, with no blank tokens. |
123
+ | `encoder_out_lens` | Int32/int64 frame counts of shape `(N,)`, each between `1` and `T`. |
124
+ | `chunk_size` | Number of vocabulary columns processed per chunk. Smaller chunks reduce temporary projection memory and can cost speed. |
125
+ | `output_beam` | Positive, finite log-score beam used for alignment pruning. Larger beams retain more alignments and use more memory. |
126
+ | `max_states` | Lattice state limit. k2 may narrow the effective beam to stay within this limit. |
127
+ | `blank_id` | Blank token ID in the original vocabulary; it is remapped internally to column zero for k2. |
128
+
129
+ The return value is a scalar `float32` **sum** over utterances. Nonfinite
130
+ per-utterance losses are excluded with a warning, including losses from
131
+ impossible target alignments. With finite inputs, impossible alignments
132
+ contribute zero loss and zero gradients. Apply any desired loss normalization
133
+ explicitly.
134
+
135
+ ### Precision and device requirements
136
+
137
+ - `encoder_out`, `weight`, and `bias` may each use `float32` or `bfloat16`.
138
+ Their values must be finite. Gradients retain the corresponding input dtypes.
139
+ `float16` is unsupported.
140
+ - Disable autocast around the loss call, even when the encoder runs under mixed
141
+ precision. The loss manages its own numerical precision.
142
+ - Put the encoder states, projection parameters, and targets on the same CPU or
143
+ CUDA device. Frame counts may also stay on the CPU.
144
+ - Target IDs must be valid vocabulary indices and must exclude `blank_id`.
145
+ - Only first-order gradients are supported; double backward is unsupported.
146
+
147
+ ## License
148
+
149
+ [MIT](https://github.com/yfyeung/PrunedCTC/blob/master/LICENSE). Copyright 2026 Shanghai Jiao Tong University.
150
+ Author: Yifan Yang.
@@ -0,0 +1,11 @@
1
+ CONTRIBUTING.md
2
+ LICENSE
3
+ MANIFEST.in
4
+ README.md
5
+ pruned_ctc.py
6
+ pyproject.toml
7
+ pruned_ctc.egg-info/PKG-INFO
8
+ pruned_ctc.egg-info/SOURCES.txt
9
+ pruned_ctc.egg-info/dependency_links.txt
10
+ pruned_ctc.egg-info/requires.txt
11
+ pruned_ctc.egg-info/top_level.txt
@@ -0,0 +1,6 @@
1
+ torch>=2.4
2
+
3
+ [dev]
4
+ build>=1.2.2
5
+ ruff>=0.9
6
+ twine>=6
@@ -0,0 +1 @@
1
+ pruned_ctc
@@ -0,0 +1,418 @@
1
+ # Copyright (c) 2026 Shanghai Jiao Tong University (author: Yifan Yang)
2
+ # SPDX-License-Identifier: MIT
3
+
4
+ """Pruned CTC with full-vocabulary normalization and chunked recomputation."""
5
+
6
+ from __future__ import annotations
7
+
8
+ import importlib
9
+ import logging
10
+ import math
11
+ from numbers import Integral, Real
12
+ from types import ModuleType
13
+ from typing import TYPE_CHECKING, Any, overload
14
+
15
+ import torch
16
+ from torch import Tensor
17
+ from torch.autograd.function import once_differentiable
18
+
19
+ if TYPE_CHECKING:
20
+ import k2
21
+
22
+ __version__ = "0.1.0"
23
+ __all__ = ["pruned_ctc_loss"]
24
+
25
+ # k2 stores dense FSA labels in uint16.
26
+ MAX_REDUCED_VOCAB = 65535
27
+ DEFAULT_MAX_STATES = 100_000_000
28
+
29
+ _FLOAT_DTYPES = (torch.float32, torch.bfloat16)
30
+ _INDEX_DTYPES = (torch.int32, torch.int64)
31
+ _INT32_MAX = torch.iinfo(torch.int32).max
32
+ _LOGGER = logging.getLogger(__name__)
33
+ _warned_near_cap = False
34
+
35
+
36
+ def _load_k2() -> ModuleType:
37
+ try:
38
+ return importlib.import_module("k2")
39
+ except (ImportError, OSError) as error:
40
+ raise ImportError(
41
+ "pruned_ctc_loss requires k2 built for your PyTorch, Python, and CUDA "
42
+ "versions. Install a matching k2 build before using the loss: "
43
+ "https://k2-fsa.github.io/k2/installation/index.html"
44
+ ) from error
45
+
46
+
47
+ def _positive_integer(name: str, value: int) -> int:
48
+ if isinstance(value, bool) or not isinstance(value, Integral):
49
+ raise TypeError(f"{name} must be an integer, got {type(value).__name__}.")
50
+ if value <= 0:
51
+ raise ValueError(f"{name} must be positive, got {value}.")
52
+ return int(value)
53
+
54
+
55
+ def _validate_inputs(
56
+ encoder_out: Tensor,
57
+ weight: Tensor,
58
+ bias: Tensor | None,
59
+ y: k2.RaggedTensor,
60
+ encoder_out_lens: Tensor,
61
+ blank_id: int,
62
+ backend: ModuleType,
63
+ ) -> Tensor:
64
+ for name, value in (
65
+ ("encoder_out", encoder_out),
66
+ ("weight", weight),
67
+ ("encoder_out_lens", encoder_out_lens),
68
+ ):
69
+ if not isinstance(value, Tensor):
70
+ raise TypeError(f"{name} must be a torch.Tensor.")
71
+ if bias is not None and not isinstance(bias, Tensor):
72
+ raise TypeError("bias must be a torch.Tensor or None.")
73
+ if encoder_out.ndim != 3 or any(size == 0 for size in encoder_out.shape):
74
+ raise ValueError("encoder_out must have nonempty shape (N, T, D).")
75
+ if weight.ndim != 2 or weight.shape[0] == 0:
76
+ raise ValueError("weight must have nonempty shape (V, D).")
77
+ batch_size, max_frames, hidden_dim = encoder_out.shape
78
+ vocab_size = weight.shape[0]
79
+ if weight.shape[1] != hidden_dim:
80
+ raise ValueError("weight and encoder_out must have the same hidden dimension.")
81
+ if vocab_size > _INT32_MAX:
82
+ raise ValueError("The vocabulary must fit in k2's signed int32 label indices.")
83
+ if bias is not None and bias.shape != (vocab_size,):
84
+ raise ValueError(f"bias must have shape ({vocab_size},).")
85
+ device = encoder_out.device
86
+ if device.type not in ("cpu", "cuda"):
87
+ raise ValueError("pruned_ctc_loss supports CPU and CUDA tensors.")
88
+ for name, value in (("encoder_out", encoder_out), ("weight", weight), ("bias", bias)):
89
+ if value is None:
90
+ continue
91
+ if value.dtype not in _FLOAT_DTYPES:
92
+ raise TypeError(f"{name} must use float32 or bfloat16, got {value.dtype}.")
93
+ if value.device != device:
94
+ raise ValueError(f"{name} must be on {device}, got {value.device}.")
95
+ if torch.is_autocast_enabled(device.type):
96
+ raise RuntimeError(
97
+ "pruned_ctc_loss controls its own precision. Call it inside "
98
+ f"torch.autocast('{device.type}', enabled=False)."
99
+ )
100
+ if isinstance(blank_id, bool) or not isinstance(blank_id, Integral):
101
+ raise TypeError("blank_id must be an integer.")
102
+ if not 0 <= blank_id < vocab_size:
103
+ raise ValueError(f"blank_id must be in [0, {vocab_size}), got {blank_id}.")
104
+ if encoder_out_lens.ndim != 1 or encoder_out_lens.shape[0] != batch_size:
105
+ raise ValueError(f"encoder_out_lens must have shape ({batch_size},).")
106
+ if encoder_out_lens.dtype not in _INDEX_DTYPES:
107
+ raise TypeError("encoder_out_lens must use int32 or int64.")
108
+ if encoder_out_lens.device.type != "cpu" and encoder_out_lens.device != device:
109
+ raise ValueError("encoder_out_lens must be on CPU or the encoder device.")
110
+ lengths_cpu = encoder_out_lens.detach().to(device="cpu", dtype=torch.int64)
111
+ lengths_cpu = lengths_cpu.clone(memory_format=torch.contiguous_format)
112
+ if int(lengths_cpu.min()) < 1 or int(lengths_cpu.max()) > max_frames:
113
+ raise ValueError(f"encoder_out_lens must contain values in [1, {max_frames}].")
114
+ if int(lengths_cpu.sum()) > _INT32_MAX:
115
+ raise ValueError("The number of valid frames exceeds k2's int32 index range.")
116
+ if not isinstance(y, backend.RaggedTensor):
117
+ raise TypeError("y must be a k2.RaggedTensor.")
118
+ if y.num_axes != 2 or y.dim0 != batch_size:
119
+ raise ValueError("y must have two axes and one target sequence per utterance.")
120
+ if y.device != device:
121
+ raise ValueError(f"y must be on {device}, got {y.device}.")
122
+ targets = y.values
123
+ if targets.dtype not in _INDEX_DTYPES:
124
+ raise TypeError("Target IDs must use int32 or int64.")
125
+ if targets.numel():
126
+ invalid_range, contains_blank = torch.stack(
127
+ (((targets < 0) | (targets >= vocab_size)).any(), (targets == blank_id).any())
128
+ ).tolist()
129
+ if invalid_range:
130
+ raise ValueError(f"Target IDs must be in [0, {vocab_size}).")
131
+ if contains_blank:
132
+ raise ValueError("Target sequences must not contain blank_id.")
133
+ return lengths_cpu
134
+
135
+
136
+ @overload
137
+ def _at_least_fp32(value: Tensor) -> Tensor: ...
138
+
139
+
140
+ @overload
141
+ def _at_least_fp32(value: None) -> None: ...
142
+
143
+
144
+ def _at_least_fp32(value: Tensor | None) -> Tensor | None:
145
+ if value is None or value.dtype in (torch.float32, torch.float64):
146
+ return value
147
+ return value.float()
148
+
149
+
150
+ def _chunk_logits(hidden: Tensor, weight: Tensor, bias: Tensor | None) -> Tensor:
151
+ if bias is None:
152
+ return torch.mm(hidden, weight.t())
153
+ return torch.addmm(bias, hidden, weight.t())
154
+
155
+
156
+ class _FusedReducedLogSoftmax(torch.autograd.Function):
157
+ """Recompute dense softmax gradients in chunks; support first derivatives only."""
158
+
159
+ @staticmethod
160
+ def forward(
161
+ ctx: Any,
162
+ hidden: Tensor,
163
+ weight: Tensor,
164
+ bias: Tensor | None,
165
+ selected: Tensor,
166
+ chunk_size: int = 4096,
167
+ ) -> Tensor:
168
+ hidden_dim = hidden.shape[-1]
169
+ vocab_size = weight.shape[0]
170
+ flat = _at_least_fp32(hidden.reshape(-1, hidden_dim))
171
+ num_frames = flat.shape[0]
172
+ dp_dtype = torch.float64 if flat.dtype == torch.float64 else torch.float32
173
+
174
+ # Float64 normalization limits cancellation error in selected-class gradients.
175
+ running_max = flat.new_full((num_frames,), -math.inf, dtype=torch.float64)
176
+ running_sum = flat.new_zeros((num_frames,), dtype=torch.float64)
177
+ for start in range(0, vocab_size, chunk_size):
178
+ end = min(start + chunk_size, vocab_size)
179
+ logits = _chunk_logits(
180
+ flat,
181
+ _at_least_fp32(weight[start:end]),
182
+ _at_least_fp32(None if bias is None else bias[start:end]),
183
+ ).double()
184
+ new_max = torch.maximum(running_max, logits.amax(dim=1))
185
+ running_sum.mul_((running_max - new_max).exp_()).add_(
186
+ logits.sub_(new_max.unsqueeze(1)).exp_().sum(dim=1)
187
+ )
188
+ running_max = new_max
189
+ # Release the full chunk before allocating selected-class logits.
190
+ del logits
191
+ normalizer = running_max.add_(running_sum.log_())
192
+ selected_weight = _at_least_fp32(weight.index_select(0, selected))
193
+ selected_bias = None if bias is None else _at_least_fp32(bias.index_select(0, selected))
194
+ log_probs = _chunk_logits(flat, selected_weight, selected_bias)
195
+ # Subtract in float64, then round once to the DP dtype.
196
+ log_probs = log_probs.double().sub_(normalizer.unsqueeze(1)).to(dp_dtype)
197
+ ctx.save_for_backward(hidden, weight, bias, selected, normalizer.to(dp_dtype), log_probs)
198
+ ctx.chunk_size = chunk_size
199
+ return log_probs.view(*hidden.shape[:-1], selected.numel())
200
+
201
+ @staticmethod
202
+ @once_differentiable
203
+ def backward(
204
+ ctx: Any, grad_log_probs: Tensor
205
+ ) -> tuple[Tensor | None, Tensor | None, Tensor | None, None, None]:
206
+ hidden, weight, bias, selected, normalizer, saved_log_probs = ctx.saved_tensors
207
+ # Backward can run outside the caller's forward autocast context.
208
+ with torch.autocast(hidden.device.type, enabled=False):
209
+ chunk_size = ctx.chunk_size
210
+ hidden_dim = hidden.shape[-1]
211
+ vocab_size = weight.shape[0]
212
+ need_hidden, need_weight, need_bias = ctx.needs_input_grad[:3]
213
+ need_bias = need_bias and bias is not None
214
+ flat = _at_least_fp32(hidden.reshape(-1, hidden_dim))
215
+ grad = _at_least_fp32(grad_log_probs.reshape(-1, selected.numel()))
216
+ # Preserve upstream scaling and zero gradients at omitted frames.
217
+ negative_sum = (-grad.sum(dim=1)).unsqueeze(1)
218
+ grad_hidden = torch.zeros_like(flat) if need_hidden else None
219
+ grad_weight = torch.empty_like(weight) if need_weight else None
220
+ grad_bias = torch.empty_like(bias) if need_bias else None
221
+ cast_weight = need_weight and weight.dtype not in (torch.float32, torch.float64)
222
+ cast_bias = need_bias and bias.dtype not in (torch.float32, torch.float64)
223
+ # Selected IDs need not be sorted when the original blank ID is nonzero.
224
+ chunk_ids = torch.div(selected, chunk_size, rounding_mode="floor")
225
+ for chunk_index, start in enumerate(range(0, vocab_size, chunk_size)):
226
+ end = min(start + chunk_size, vocab_size)
227
+ weight_chunk = _at_least_fp32(weight[start:end])
228
+ grad_logits = _chunk_logits(
229
+ flat, weight_chunk, _at_least_fp32(None if bias is None else bias[start:end])
230
+ )
231
+ grad_logits.sub_(normalizer.unsqueeze(1)).exp_().mul_(negative_sum)
232
+ positions = (chunk_ids == chunk_index).nonzero().flatten()
233
+ if positions.numel():
234
+ local_ids = selected.index_select(0, positions) - start
235
+ # Reuse forward probabilities and combine cancelling terms before reductions.
236
+ grad_logits.index_copy_(
237
+ 1,
238
+ local_ids,
239
+ saved_log_probs.index_select(1, positions).exp().mul_(negative_sum),
240
+ )
241
+ grad_logits.index_add_(1, local_ids, grad.index_select(1, positions))
242
+ if need_hidden:
243
+ grad_hidden.addmm_(grad_logits, weight_chunk)
244
+ if need_weight:
245
+ if cast_weight:
246
+ grad_weight[start:end].copy_(torch.mm(grad_logits.t(), flat))
247
+ else:
248
+ torch.mm(grad_logits.t(), flat, out=grad_weight[start:end])
249
+ if need_bias:
250
+ if cast_bias:
251
+ grad_bias[start:end].copy_(grad_logits.sum(dim=0))
252
+ else:
253
+ torch.sum(grad_logits, dim=0, out=grad_bias[start:end])
254
+ return (
255
+ grad_hidden.view_as(hidden).to(hidden.dtype) if need_hidden else None,
256
+ grad_weight,
257
+ grad_bias,
258
+ None,
259
+ None,
260
+ )
261
+
262
+
263
+ def pruned_ctc_loss(
264
+ encoder_out: Tensor,
265
+ weight: Tensor,
266
+ bias: Tensor | None,
267
+ y: k2.RaggedTensor,
268
+ encoder_out_lens: Tensor,
269
+ *,
270
+ chunk_size: int = 4096,
271
+ output_beam: float = 100.0,
272
+ max_states: int = DEFAULT_MAX_STATES,
273
+ blank_id: int = 0,
274
+ ) -> Tensor:
275
+ """Compute summed CTC loss with reduced vocabulary and pruned alignments.
276
+
277
+ Args:
278
+ encoder_out: Encoder states of shape (N, T, D), in float32 or bfloat16.
279
+ weight: Projection weights of shape (V, D), in float32 or bfloat16.
280
+ bias: Projection bias of shape (V,), in float32 or bfloat16, or None.
281
+ y: Two-axis k2.RaggedTensor of target IDs, excluding blank_id.
282
+ encoder_out_lens: Int32/int64 frame counts of shape (N,), each in [1, T].
283
+ chunk_size: Positive number of vocabulary columns per projection chunk.
284
+ output_beam: Positive finite log-score beam, representable in float32.
285
+ max_states: Positive lattice state limit. k2 can tighten the effective beam.
286
+ blank_id: Blank ID in the original vocabulary, remapped to column zero.
287
+
288
+ States, projection parameters, and targets must share a CPU or CUDA device.
289
+ Floating-point inputs must be finite. Lengths may remain on CPU.
290
+ Disable autocast when calling this function;
291
+ gradients retain each floating-point input's dtype. Vocabulary reduction is
292
+ exact, while finite-beam alignment pruning is approximate. k2's lattice
293
+ limits can further restrict the retained alignments.
294
+
295
+ Returns:
296
+ A float32 scalar summed over utterances. Nonfinite per-utterance losses
297
+ are excluded with a logged warning. Impossible alignments contribute
298
+ zero gradients when the floating-point inputs are finite.
299
+ Only first-order differentiation is supported.
300
+
301
+ Raises:
302
+ ImportError: A compatible k2 installation is unavailable.
303
+ TypeError: An input has an unsupported type or dtype.
304
+ ValueError: A shape, device, index, or configuration value is invalid.
305
+ RuntimeError: Autocast is enabled on the input device.
306
+ """
307
+ chunk_size = _positive_integer("chunk_size", chunk_size)
308
+ max_states = _positive_integer("max_states", max_states)
309
+ if max_states > _INT32_MAX:
310
+ raise ValueError("max_states must fit in a signed int32.")
311
+ if isinstance(output_beam, bool) or not isinstance(output_beam, Real):
312
+ raise TypeError("output_beam must be a real number.")
313
+ output_beam = float(output_beam)
314
+ beam_limits = torch.finfo(torch.float32)
315
+ if not math.isfinite(output_beam) or not (
316
+ beam_limits.tiny * beam_limits.eps <= output_beam <= beam_limits.max
317
+ ):
318
+ raise ValueError("output_beam must be positive, finite, and representable in float32.")
319
+ backend = _load_k2()
320
+ lengths_cpu = _validate_inputs(
321
+ encoder_out, weight, bias, y, encoder_out_lens, blank_id, backend
322
+ )
323
+ chunk_size = min(chunk_size, weight.shape[0])
324
+ blank_id = int(blank_id)
325
+ device = encoder_out.device
326
+ batch_size, max_frames, hidden_dim = encoder_out.shape
327
+ targets = y.values.to(torch.int32)
328
+ distinct = torch.unique(targets, sorted=True)
329
+ selected = torch.cat((targets.new_full((1,), blank_id), distinct))
330
+ if selected.numel() > MAX_REDUCED_VOCAB:
331
+ if batch_size == 1:
332
+ raise ValueError(
333
+ f"One utterance needs {selected.numel()} reduced classes; "
334
+ f"k2 supports at most {MAX_REDUCED_VOCAB}. Shorten its transcript."
335
+ )
336
+ midpoint = batch_size // 2
337
+ _LOGGER.warning(
338
+ "Reduced vocabulary has %d classes; splitting %d utterances to fit k2's limit.",
339
+ selected.numel(),
340
+ batch_size,
341
+ )
342
+ return sum(
343
+ pruned_ctc_loss(
344
+ encoder_out=encoder_out[start:end],
345
+ weight=weight,
346
+ bias=bias,
347
+ y=y[start:end],
348
+ encoder_out_lens=encoder_out_lens[start:end],
349
+ chunk_size=chunk_size,
350
+ output_beam=output_beam,
351
+ max_states=max_states,
352
+ blank_id=blank_id,
353
+ )
354
+ for start, end in ((0, midpoint), (midpoint, batch_size))
355
+ )
356
+ remapped = (torch.searchsorted(distinct, targets) + 1).to(torch.int32)
357
+ num_valid_frames = int(lengths_cpu.sum())
358
+ packed = num_valid_frames < batch_size * max_frames
359
+ if packed:
360
+ lengths_device = lengths_cpu.to(device=device)
361
+ starts_device = torch.cumsum(lengths_device, dim=0) - lengths_device
362
+ utterance_ids = torch.repeat_interleave(
363
+ torch.arange(batch_size, device=device), lengths_device, output_size=num_valid_frames
364
+ )
365
+ frame_ids = torch.arange(num_valid_frames, device=device) - starts_device[utterance_ids]
366
+ hidden = encoder_out.reshape(-1, hidden_dim).index_select(
367
+ 0, utterance_ids * max_frames + frame_ids
368
+ )
369
+ starts_cpu = torch.cumsum(lengths_cpu, dim=0) - lengths_cpu
370
+ sequence_ids_cpu = torch.zeros(batch_size, dtype=torch.int64, device="cpu")
371
+ else:
372
+ hidden = encoder_out
373
+ starts_cpu = torch.zeros(batch_size, dtype=torch.int64, device="cpu")
374
+ sequence_ids_cpu = torch.arange(batch_size, dtype=torch.int64, device="cpu")
375
+ log_probs = _FusedReducedLogSoftmax.apply(hidden, weight, bias, selected.long(), chunk_size)
376
+ if packed:
377
+ log_probs = log_probs.unsqueeze(0)
378
+ # k2 creates CPU supervision indices internally, regardless of the caller's default device.
379
+ with torch.device("cpu"):
380
+ decoding_graph = backend.ctc_graph(backend.RaggedTensor(y.shape, remapped), modified=False)
381
+ # A contiguous length vector avoids stride issues for a single utterance.
382
+ # Apply the same stable permutation to segments and graphs.
383
+ order = torch.argsort(lengths_cpu, descending=True, stable=True)
384
+ segments = torch.stack(
385
+ (sequence_ids_cpu[order], starts_cpu[order], lengths_cpu[order]), dim=1
386
+ ).to(torch.int32)
387
+ decoding_graph = backend.index_fsa(
388
+ decoding_graph, order.to(device=device, dtype=torch.int32)
389
+ )
390
+ dense_fsa = backend.DenseFsaVec(log_probs, segments)
391
+ lattice = backend.intersect_dense(
392
+ a_fsas=decoding_graph,
393
+ b_fsas=dense_fsa,
394
+ output_beam=output_beam,
395
+ max_states=max_states,
396
+ )
397
+ losses = -lattice.get_tot_scores(log_semiring=True, use_double_scores=True).float()
398
+ global _warned_near_cap
399
+ num_states = lattice.arcs.tot_size(1)
400
+ if num_states > 0.5 * max_states and not _warned_near_cap:
401
+ _warned_near_cap = True
402
+ _LOGGER.warning(
403
+ "Output lattice uses %d states (max_states=%d). k2 may tighten the beam "
404
+ "when lattice limits are reached; this warning is a heuristic.",
405
+ num_states,
406
+ max_states,
407
+ )
408
+ loss = losses.sum()
409
+ if bool(torch.isfinite(loss)):
410
+ return loss
411
+ finite = torch.isfinite(losses)
412
+ invalid_indices = order[~finite.cpu()].tolist()
413
+ _LOGGER.warning(
414
+ "Dropping %d utterances with nonfinite CTC loss; batch indices: %s.",
415
+ len(invalid_indices),
416
+ invalid_indices,
417
+ )
418
+ return torch.where(finite, losses, losses.new_zeros(())).sum()
@@ -0,0 +1,43 @@
1
+ [build-system]
2
+ requires = ["setuptools>=77.0.3"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "pruned-ctc"
7
+ dynamic = ["version"]
8
+ description = "Memory-efficient CTC loss with exact vocabulary reduction and pruned alignments"
9
+ readme = "README.md"
10
+ requires-python = ">=3.10"
11
+ license = "MIT"
12
+ license-files = ["LICENSE"]
13
+ authors = [{ name = "Yifan Yang" }]
14
+ keywords = ["ctc", "speech-recognition", "pytorch", "sequence-modeling"]
15
+ classifiers = [
16
+ "Development Status :: 4 - Beta",
17
+ "Intended Audience :: Science/Research",
18
+ "Programming Language :: Python :: 3",
19
+ "Programming Language :: Python :: 3 :: Only",
20
+ "Topic :: Scientific/Engineering :: Artificial Intelligence",
21
+ ]
22
+ dependencies = ["torch>=2.4"]
23
+
24
+ [project.optional-dependencies]
25
+ dev = ["build>=1.2.2", "ruff>=0.9", "twine>=6"]
26
+
27
+ [project.urls]
28
+ Homepage = "https://github.com/yfyeung/PrunedCTC"
29
+ Repository = "https://github.com/yfyeung/PrunedCTC"
30
+ Issues = "https://github.com/yfyeung/PrunedCTC/issues"
31
+
32
+ [tool.setuptools]
33
+ py-modules = ["pruned_ctc"]
34
+
35
+ [tool.setuptools.dynamic]
36
+ version = { attr = "pruned_ctc.__version__" }
37
+
38
+ [tool.ruff]
39
+ target-version = "py310"
40
+ line-length = 100
41
+
42
+ [tool.ruff.lint]
43
+ select = ["E", "F", "I", "B", "UP"]
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+