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.
- pruned_ctc-0.1.0/CONTRIBUTING.md +59 -0
- pruned_ctc-0.1.0/LICENSE +21 -0
- pruned_ctc-0.1.0/MANIFEST.in +1 -0
- pruned_ctc-0.1.0/PKG-INFO +150 -0
- pruned_ctc-0.1.0/README.md +125 -0
- pruned_ctc-0.1.0/pruned_ctc.egg-info/PKG-INFO +150 -0
- pruned_ctc-0.1.0/pruned_ctc.egg-info/SOURCES.txt +11 -0
- pruned_ctc-0.1.0/pruned_ctc.egg-info/dependency_links.txt +1 -0
- pruned_ctc-0.1.0/pruned_ctc.egg-info/requires.txt +6 -0
- pruned_ctc-0.1.0/pruned_ctc.egg-info/top_level.txt +1 -0
- pruned_ctc-0.1.0/pruned_ctc.py +418 -0
- pruned_ctc-0.1.0/pyproject.toml +43 -0
- pruned_ctc-0.1.0/setup.cfg +4 -0
|
@@ -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.
|
pruned_ctc-0.1.0/LICENSE
ADDED
|
@@ -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 @@
|
|
|
1
|
+
|
|
@@ -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"]
|