Feature Extraction
Transformers
Safetensors
sheetsage2
audio
music
music-transcription
midi
abc-notation
custom_code
Instructions to use rAVEUK/SheetSage2 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use rAVEUK/SheetSage2 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("feature-extraction", model="rAVEUK/SheetSage2", trust_remote_code=True)# pip install -U transformers accelerate # Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("rAVEUK/SheetSage2", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
Duplicate from m-a-p/SheetSage2
Browse filesCo-authored-by: Ruibin Yuan <a43992899@users.noreply.huggingface.co>
This view is limited to 50 files because it contains too many changes. See raw diff
- .gitattributes +3 -0
- LICENSE +424 -0
- README.md +201 -0
- THIRD_PARTY_NOTICES.md +8 -0
- __init__.py +1 -0
- assets/architecture.png +3 -0
- audio_sheetsage2.py +92 -0
- benchmark_results.json +379 -0
- config.json +75 -0
- configuration_mert2.py +76 -0
- configuration_sheetsage2.py +88 -0
- durations_sheetsage2.py +14 -0
- exports_sheetsage2.py +200 -0
- generation_sheetsage2.py +634 -0
- infer.py +78 -0
- io_sheetsage2.py +67 -0
- labels_sheetsage2.py +29 -0
- midi_sheetsage2.py +92 -0
- model.safetensors +3 -0
- modeling_mert2.py +361 -0
- modeling_sheetsage2.py +448 -0
- notation_sheetsage2.py +1570 -0
- pipeline_sheetsage2.py +259 -0
- processing_sheetsage2.py +109 -0
- processor_config.json +11 -0
- render.py +45 -0
- render_assets/DejaVuSans.ttf +3 -0
- render_assets/LICENSE.abcjs +21 -0
- render_assets/LICENSE.font +187 -0
- render_assets/abcjs-basic-min.js +0 -0
- render_assets/manifest.json +106 -0
- render_assets/renderer.js +165 -0
- render_assets/soundfonts/ATTRIBUTION.md +11 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A0.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A1.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A2.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A3.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A4.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A5.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A6.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/A7.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/Ab1.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/Ab2.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/Ab3.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/Ab4.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/Ab5.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/Ab6.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/Ab7.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/B0.mp3 +0 -0
- render_assets/soundfonts/acoustic_grand_piano-mp3/B1.mp3 +0 -0
.gitattributes
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.png filter=lfs diff=lfs merge=lfs -text
|
| 3 |
+
*.ttf filter=lfs diff=lfs merge=lfs -text
|
LICENSE
ADDED
|
@@ -0,0 +1,424 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
MERT2 model-weight license
|
| 2 |
+
|
| 3 |
+
The MERT2-30s and MERT2-FS checkpoint weights in model.safetensors are
|
| 4 |
+
licensed under Creative Commons Attribution-NonCommercial 4.0 International
|
| 5 |
+
(CC BY-NC 4.0): https://creativecommons.org/licenses/by-nc/4.0/
|
| 6 |
+
|
| 7 |
+
For attribution, identify MERT2, the model name, and its source repository:
|
| 8 |
+
https://huggingface.co/m-a-p/MERT-v2-30s
|
| 9 |
+
https://huggingface.co/m-a-p/MERT-v2-FullSong
|
| 10 |
+
|
| 11 |
+
This weight license does not replace separately applicable licenses for code
|
| 12 |
+
or dependencies. See THIRD_PARTY_NOTICES.md.
|
| 13 |
+
|
| 14 |
+
The official license text follows without modification.
|
| 15 |
+
Source: https://creativecommons.org/licenses/by-nc/4.0/legalcode.txt
|
| 16 |
+
|
| 17 |
+
Attribution-NonCommercial 4.0 International
|
| 18 |
+
|
| 19 |
+
=======================================================================
|
| 20 |
+
|
| 21 |
+
Creative Commons Corporation ("Creative Commons") is not a law firm and
|
| 22 |
+
does not provide legal services or legal advice. Distribution of
|
| 23 |
+
Creative Commons public licenses does not create a lawyer-client or
|
| 24 |
+
other relationship. Creative Commons makes its licenses and related
|
| 25 |
+
information available on an "as-is" basis. Creative Commons gives no
|
| 26 |
+
warranties regarding its licenses, any material licensed under their
|
| 27 |
+
terms and conditions, or any related information. Creative Commons
|
| 28 |
+
disclaims all liability for damages resulting from their use to the
|
| 29 |
+
fullest extent possible.
|
| 30 |
+
|
| 31 |
+
Using Creative Commons Public Licenses
|
| 32 |
+
|
| 33 |
+
Creative Commons public licenses provide a standard set of terms and
|
| 34 |
+
conditions that creators and other rights holders may use to share
|
| 35 |
+
original works of authorship and other material subject to copyright
|
| 36 |
+
and certain other rights specified in the public license below. The
|
| 37 |
+
following considerations are for informational purposes only, are not
|
| 38 |
+
exhaustive, and do not form part of our licenses.
|
| 39 |
+
|
| 40 |
+
Considerations for licensors: Our public licenses are
|
| 41 |
+
intended for use by those authorized to give the public
|
| 42 |
+
permission to use material in ways otherwise restricted by
|
| 43 |
+
copyright and certain other rights. Our licenses are
|
| 44 |
+
irrevocable. Licensors should read and understand the terms
|
| 45 |
+
and conditions of the license they choose before applying it.
|
| 46 |
+
Licensors should also secure all rights necessary before
|
| 47 |
+
applying our licenses so that the public can reuse the
|
| 48 |
+
material as expected. Licensors should clearly mark any
|
| 49 |
+
material not subject to the license. This includes other CC-
|
| 50 |
+
licensed material, or material used under an exception or
|
| 51 |
+
limitation to copyright. More considerations for licensors:
|
| 52 |
+
wiki.creativecommons.org/Considerations_for_licensors
|
| 53 |
+
|
| 54 |
+
Considerations for the public: By using one of our public
|
| 55 |
+
licenses, a licensor grants the public permission to use the
|
| 56 |
+
licensed material under specified terms and conditions. If
|
| 57 |
+
the licensor's permission is not necessary for any reason--for
|
| 58 |
+
example, because of any applicable exception or limitation to
|
| 59 |
+
copyright--then that use is not regulated by the license. Our
|
| 60 |
+
licenses grant only permissions under copyright and certain
|
| 61 |
+
other rights that a licensor has authority to grant. Use of
|
| 62 |
+
the licensed material may still be restricted for other
|
| 63 |
+
reasons, including because others have copyright or other
|
| 64 |
+
rights in the material. A licensor may make special requests,
|
| 65 |
+
such as asking that all changes be marked or described.
|
| 66 |
+
Although not required by our licenses, you are encouraged to
|
| 67 |
+
respect those requests where reasonable. More considerations
|
| 68 |
+
for the public:
|
| 69 |
+
wiki.creativecommons.org/Considerations_for_licensees
|
| 70 |
+
|
| 71 |
+
=======================================================================
|
| 72 |
+
|
| 73 |
+
Creative Commons Attribution-NonCommercial 4.0 International Public
|
| 74 |
+
License
|
| 75 |
+
|
| 76 |
+
By exercising the Licensed Rights (defined below), You accept and agree
|
| 77 |
+
to be bound by the terms and conditions of this Creative Commons
|
| 78 |
+
Attribution-NonCommercial 4.0 International Public License ("Public
|
| 79 |
+
License"). To the extent this Public License may be interpreted as a
|
| 80 |
+
contract, You are granted the Licensed Rights in consideration of Your
|
| 81 |
+
acceptance of these terms and conditions, and the Licensor grants You
|
| 82 |
+
such rights in consideration of benefits the Licensor receives from
|
| 83 |
+
making the Licensed Material available under these terms and
|
| 84 |
+
conditions.
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
Section 1 -- Definitions.
|
| 88 |
+
|
| 89 |
+
a. Adapted Material means material subject to Copyright and Similar
|
| 90 |
+
Rights that is derived from or based upon the Licensed Material
|
| 91 |
+
and in which the Licensed Material is translated, altered,
|
| 92 |
+
arranged, transformed, or otherwise modified in a manner requiring
|
| 93 |
+
permission under the Copyright and Similar Rights held by the
|
| 94 |
+
Licensor. For purposes of this Public License, where the Licensed
|
| 95 |
+
Material is a musical work, performance, or sound recording,
|
| 96 |
+
Adapted Material is always produced where the Licensed Material is
|
| 97 |
+
synched in timed relation with a moving image.
|
| 98 |
+
|
| 99 |
+
b. Adapter's License means the license You apply to Your Copyright
|
| 100 |
+
and Similar Rights in Your contributions to Adapted Material in
|
| 101 |
+
accordance with the terms and conditions of this Public License.
|
| 102 |
+
|
| 103 |
+
c. Copyright and Similar Rights means copyright and/or similar rights
|
| 104 |
+
closely related to copyright including, without limitation,
|
| 105 |
+
performance, broadcast, sound recording, and Sui Generis Database
|
| 106 |
+
Rights, without regard to how the rights are labeled or
|
| 107 |
+
categorized. For purposes of this Public License, the rights
|
| 108 |
+
specified in Section 2(b)(1)-(2) are not Copyright and Similar
|
| 109 |
+
Rights.
|
| 110 |
+
d. Effective Technological Measures means those measures that, in the
|
| 111 |
+
absence of proper authority, may not be circumvented under laws
|
| 112 |
+
fulfilling obligations under Article 11 of the WIPO Copyright
|
| 113 |
+
Treaty adopted on December 20, 1996, and/or similar international
|
| 114 |
+
agreements.
|
| 115 |
+
|
| 116 |
+
e. Exceptions and Limitations means fair use, fair dealing, and/or
|
| 117 |
+
any other exception or limitation to Copyright and Similar Rights
|
| 118 |
+
that applies to Your use of the Licensed Material.
|
| 119 |
+
|
| 120 |
+
f. Licensed Material means the artistic or literary work, database,
|
| 121 |
+
or other material to which the Licensor applied this Public
|
| 122 |
+
License.
|
| 123 |
+
|
| 124 |
+
g. Licensed Rights means the rights granted to You subject to the
|
| 125 |
+
terms and conditions of this Public License, which are limited to
|
| 126 |
+
all Copyright and Similar Rights that apply to Your use of the
|
| 127 |
+
Licensed Material and that the Licensor has authority to license.
|
| 128 |
+
|
| 129 |
+
h. Licensor means the individual(s) or entity(ies) granting rights
|
| 130 |
+
under this Public License.
|
| 131 |
+
|
| 132 |
+
i. NonCommercial means not primarily intended for or directed towards
|
| 133 |
+
commercial advantage or monetary compensation. For purposes of
|
| 134 |
+
this Public License, the exchange of the Licensed Material for
|
| 135 |
+
other material subject to Copyright and Similar Rights by digital
|
| 136 |
+
file-sharing or similar means is NonCommercial provided there is
|
| 137 |
+
no payment of monetary compensation in connection with the
|
| 138 |
+
exchange.
|
| 139 |
+
|
| 140 |
+
j. Share means to provide material to the public by any means or
|
| 141 |
+
process that requires permission under the Licensed Rights, such
|
| 142 |
+
as reproduction, public display, public performance, distribution,
|
| 143 |
+
dissemination, communication, or importation, and to make material
|
| 144 |
+
available to the public including in ways that members of the
|
| 145 |
+
public may access the material from a place and at a time
|
| 146 |
+
individually chosen by them.
|
| 147 |
+
|
| 148 |
+
k. Sui Generis Database Rights means rights other than copyright
|
| 149 |
+
resulting from Directive 96/9/EC of the European Parliament and of
|
| 150 |
+
the Council of 11 March 1996 on the legal protection of databases,
|
| 151 |
+
as amended and/or succeeded, as well as other essentially
|
| 152 |
+
equivalent rights anywhere in the world.
|
| 153 |
+
|
| 154 |
+
l. You means the individual or entity exercising the Licensed Rights
|
| 155 |
+
under this Public License. Your has a corresponding meaning.
|
| 156 |
+
|
| 157 |
+
|
| 158 |
+
Section 2 -- Scope.
|
| 159 |
+
|
| 160 |
+
a. License grant.
|
| 161 |
+
|
| 162 |
+
1. Subject to the terms and conditions of this Public License,
|
| 163 |
+
the Licensor hereby grants You a worldwide, royalty-free,
|
| 164 |
+
non-sublicensable, non-exclusive, irrevocable license to
|
| 165 |
+
exercise the Licensed Rights in the Licensed Material to:
|
| 166 |
+
|
| 167 |
+
a. reproduce and Share the Licensed Material, in whole or
|
| 168 |
+
in part, for NonCommercial purposes only; and
|
| 169 |
+
|
| 170 |
+
b. produce, reproduce, and Share Adapted Material for
|
| 171 |
+
NonCommercial purposes only.
|
| 172 |
+
|
| 173 |
+
2. Exceptions and Limitations. For the avoidance of doubt, where
|
| 174 |
+
Exceptions and Limitations apply to Your use, this Public
|
| 175 |
+
License does not apply, and You do not need to comply with
|
| 176 |
+
its terms and conditions.
|
| 177 |
+
|
| 178 |
+
3. Term. The term of this Public License is specified in Section
|
| 179 |
+
6(a).
|
| 180 |
+
|
| 181 |
+
4. Media and formats; technical modifications allowed. The
|
| 182 |
+
Licensor authorizes You to exercise the Licensed Rights in
|
| 183 |
+
all media and formats whether now known or hereafter created,
|
| 184 |
+
and to make technical modifications necessary to do so. The
|
| 185 |
+
Licensor waives and/or agrees not to assert any right or
|
| 186 |
+
authority to forbid You from making technical modifications
|
| 187 |
+
necessary to exercise the Licensed Rights, including
|
| 188 |
+
technical modifications necessary to circumvent Effective
|
| 189 |
+
Technological Measures. For purposes of this Public License,
|
| 190 |
+
simply making modifications authorized by this Section 2(a)
|
| 191 |
+
(4) never produces Adapted Material.
|
| 192 |
+
|
| 193 |
+
5. Downstream recipients.
|
| 194 |
+
|
| 195 |
+
a. Offer from the Licensor -- Licensed Material. Every
|
| 196 |
+
recipient of the Licensed Material automatically
|
| 197 |
+
receives an offer from the Licensor to exercise the
|
| 198 |
+
Licensed Rights under the terms and conditions of this
|
| 199 |
+
Public License.
|
| 200 |
+
|
| 201 |
+
b. No downstream restrictions. You may not offer or impose
|
| 202 |
+
any additional or different terms or conditions on, or
|
| 203 |
+
apply any Effective Technological Measures to, the
|
| 204 |
+
Licensed Material if doing so restricts exercise of the
|
| 205 |
+
Licensed Rights by any recipient of the Licensed
|
| 206 |
+
Material.
|
| 207 |
+
|
| 208 |
+
6. No endorsement. Nothing in this Public License constitutes or
|
| 209 |
+
may be construed as permission to assert or imply that You
|
| 210 |
+
are, or that Your use of the Licensed Material is, connected
|
| 211 |
+
with, or sponsored, endorsed, or granted official status by,
|
| 212 |
+
the Licensor or others designated to receive attribution as
|
| 213 |
+
provided in Section 3(a)(1)(A)(i).
|
| 214 |
+
|
| 215 |
+
b. Other rights.
|
| 216 |
+
|
| 217 |
+
1. Moral rights, such as the right of integrity, are not
|
| 218 |
+
licensed under this Public License, nor are publicity,
|
| 219 |
+
privacy, and/or other similar personality rights; however, to
|
| 220 |
+
the extent possible, the Licensor waives and/or agrees not to
|
| 221 |
+
assert any such rights held by the Licensor to the limited
|
| 222 |
+
extent necessary to allow You to exercise the Licensed
|
| 223 |
+
Rights, but not otherwise.
|
| 224 |
+
|
| 225 |
+
2. Patent and trademark rights are not licensed under this
|
| 226 |
+
Public License.
|
| 227 |
+
|
| 228 |
+
3. To the extent possible, the Licensor waives any right to
|
| 229 |
+
collect royalties from You for the exercise of the Licensed
|
| 230 |
+
Rights, whether directly or through a collecting society
|
| 231 |
+
under any voluntary or waivable statutory or compulsory
|
| 232 |
+
licensing scheme. In all other cases the Licensor expressly
|
| 233 |
+
reserves any right to collect such royalties, including when
|
| 234 |
+
the Licensed Material is used other than for NonCommercial
|
| 235 |
+
purposes.
|
| 236 |
+
|
| 237 |
+
|
| 238 |
+
Section 3 -- License Conditions.
|
| 239 |
+
|
| 240 |
+
Your exercise of the Licensed Rights is expressly made subject to the
|
| 241 |
+
following conditions.
|
| 242 |
+
|
| 243 |
+
a. Attribution.
|
| 244 |
+
|
| 245 |
+
1. If You Share the Licensed Material (including in modified
|
| 246 |
+
form), You must:
|
| 247 |
+
|
| 248 |
+
a. retain the following if it is supplied by the Licensor
|
| 249 |
+
with the Licensed Material:
|
| 250 |
+
|
| 251 |
+
i. identification of the creator(s) of the Licensed
|
| 252 |
+
Material and any others designated to receive
|
| 253 |
+
attribution, in any reasonable manner requested by
|
| 254 |
+
the Licensor (including by pseudonym if
|
| 255 |
+
designated);
|
| 256 |
+
|
| 257 |
+
ii. a copyright notice;
|
| 258 |
+
|
| 259 |
+
iii. a notice that refers to this Public License;
|
| 260 |
+
|
| 261 |
+
iv. a notice that refers to the disclaimer of
|
| 262 |
+
warranties;
|
| 263 |
+
|
| 264 |
+
v. a URI or hyperlink to the Licensed Material to the
|
| 265 |
+
extent reasonably practicable;
|
| 266 |
+
|
| 267 |
+
b. indicate if You modified the Licensed Material and
|
| 268 |
+
retain an indication of any previous modifications; and
|
| 269 |
+
|
| 270 |
+
c. indicate the Licensed Material is licensed under this
|
| 271 |
+
Public License, and include the text of, or the URI or
|
| 272 |
+
hyperlink to, this Public License.
|
| 273 |
+
|
| 274 |
+
2. You may satisfy the conditions in Section 3(a)(1) in any
|
| 275 |
+
reasonable manner based on the medium, means, and context in
|
| 276 |
+
which You Share the Licensed Material. For example, it may be
|
| 277 |
+
reasonable to satisfy the conditions by providing a URI or
|
| 278 |
+
hyperlink to a resource that includes the required
|
| 279 |
+
information.
|
| 280 |
+
|
| 281 |
+
3. If requested by the Licensor, You must remove any of the
|
| 282 |
+
information required by Section 3(a)(1)(A) to the extent
|
| 283 |
+
reasonably practicable.
|
| 284 |
+
|
| 285 |
+
4. If You Share Adapted Material You produce, the Adapter's
|
| 286 |
+
License You apply must not prevent recipients of the Adapted
|
| 287 |
+
Material from complying with this Public License.
|
| 288 |
+
|
| 289 |
+
|
| 290 |
+
Section 4 -- Sui Generis Database Rights.
|
| 291 |
+
|
| 292 |
+
Where the Licensed Rights include Sui Generis Database Rights that
|
| 293 |
+
apply to Your use of the Licensed Material:
|
| 294 |
+
|
| 295 |
+
a. for the avoidance of doubt, Section 2(a)(1) grants You the right
|
| 296 |
+
to extract, reuse, reproduce, and Share all or a substantial
|
| 297 |
+
portion of the contents of the database for NonCommercial purposes
|
| 298 |
+
only;
|
| 299 |
+
|
| 300 |
+
b. if You include all or a substantial portion of the database
|
| 301 |
+
contents in a database in which You have Sui Generis Database
|
| 302 |
+
Rights, then the database in which You have Sui Generis Database
|
| 303 |
+
Rights (but not its individual contents) is Adapted Material; and
|
| 304 |
+
|
| 305 |
+
c. You must comply with the conditions in Section 3(a) if You Share
|
| 306 |
+
all or a substantial portion of the contents of the database.
|
| 307 |
+
|
| 308 |
+
For the avoidance of doubt, this Section 4 supplements and does not
|
| 309 |
+
replace Your obligations under this Public License where the Licensed
|
| 310 |
+
Rights include other Copyright and Similar Rights.
|
| 311 |
+
|
| 312 |
+
|
| 313 |
+
Section 5 -- Disclaimer of Warranties and Limitation of Liability.
|
| 314 |
+
|
| 315 |
+
a. UNLESS OTHERWISE SEPARATELY UNDERTAKEN BY THE LICENSOR, TO THE
|
| 316 |
+
EXTENT POSSIBLE, THE LICENSOR OFFERS THE LICENSED MATERIAL AS-IS
|
| 317 |
+
AND AS-AVAILABLE, AND MAKES NO REPRESENTATIONS OR WARRANTIES OF
|
| 318 |
+
ANY KIND CONCERNING THE LICENSED MATERIAL, WHETHER EXPRESS,
|
| 319 |
+
IMPLIED, STATUTORY, OR OTHER. THIS INCLUDES, WITHOUT LIMITATION,
|
| 320 |
+
WARRANTIES OF TITLE, MERCHANTABILITY, FITNESS FOR A PARTICULAR
|
| 321 |
+
PURPOSE, NON-INFRINGEMENT, ABSENCE OF LATENT OR OTHER DEFECTS,
|
| 322 |
+
ACCURACY, OR THE PRESENCE OR ABSENCE OF ERRORS, WHETHER OR NOT
|
| 323 |
+
KNOWN OR DISCOVERABLE. WHERE DISCLAIMERS OF WARRANTIES ARE NOT
|
| 324 |
+
ALLOWED IN FULL OR IN PART, THIS DISCLAIMER MAY NOT APPLY TO YOU.
|
| 325 |
+
|
| 326 |
+
b. TO THE EXTENT POSSIBLE, IN NO EVENT WILL THE LICENSOR BE LIABLE
|
| 327 |
+
TO YOU ON ANY LEGAL THEORY (INCLUDING, WITHOUT LIMITATION,
|
| 328 |
+
NEGLIGENCE) OR OTHERWISE FOR ANY DIRECT, SPECIAL, INDIRECT,
|
| 329 |
+
INCIDENTAL, CONSEQUENTIAL, PUNITIVE, EXEMPLARY, OR OTHER LOSSES,
|
| 330 |
+
COSTS, EXPENSES, OR DAMAGES ARISING OUT OF THIS PUBLIC LICENSE OR
|
| 331 |
+
USE OF THE LICENSED MATERIAL, EVEN IF THE LICENSOR HAS BEEN
|
| 332 |
+
ADVISED OF THE POSSIBILITY OF SUCH LOSSES, COSTS, EXPENSES, OR
|
| 333 |
+
DAMAGES. WHERE A LIMITATION OF LIABILITY IS NOT ALLOWED IN FULL OR
|
| 334 |
+
IN PART, THIS LIMITATION MAY NOT APPLY TO YOU.
|
| 335 |
+
|
| 336 |
+
c. The disclaimer of warranties and limitation of liability provided
|
| 337 |
+
above shall be interpreted in a manner that, to the extent
|
| 338 |
+
possible, most closely approximates an absolute disclaimer and
|
| 339 |
+
waiver of all liability.
|
| 340 |
+
|
| 341 |
+
|
| 342 |
+
Section 6 -- Term and Termination.
|
| 343 |
+
|
| 344 |
+
a. This Public License applies for the term of the Copyright and
|
| 345 |
+
Similar Rights licensed here. However, if You fail to comply with
|
| 346 |
+
this Public License, then Your rights under this Public License
|
| 347 |
+
terminate automatically.
|
| 348 |
+
|
| 349 |
+
b. Where Your right to use the Licensed Material has terminated under
|
| 350 |
+
Section 6(a), it reinstates:
|
| 351 |
+
|
| 352 |
+
1. automatically as of the date the violation is cured, provided
|
| 353 |
+
it is cured within 30 days of Your discovery of the
|
| 354 |
+
violation; or
|
| 355 |
+
|
| 356 |
+
2. upon express reinstatement by the Licensor.
|
| 357 |
+
|
| 358 |
+
For the avoidance of doubt, this Section 6(b) does not affect any
|
| 359 |
+
right the Licensor may have to seek remedies for Your violations
|
| 360 |
+
of this Public License.
|
| 361 |
+
|
| 362 |
+
c. For the avoidance of doubt, the Licensor may also offer the
|
| 363 |
+
Licensed Material under separate terms or conditions or stop
|
| 364 |
+
distributing the Licensed Material at any time; however, doing so
|
| 365 |
+
will not terminate this Public License.
|
| 366 |
+
|
| 367 |
+
d. Sections 1, 5, 6, 7, and 8 survive termination of this Public
|
| 368 |
+
License.
|
| 369 |
+
|
| 370 |
+
|
| 371 |
+
Section 7 -- Other Terms and Conditions.
|
| 372 |
+
|
| 373 |
+
a. The Licensor shall not be bound by any additional or different
|
| 374 |
+
terms or conditions communicated by You unless expressly agreed.
|
| 375 |
+
|
| 376 |
+
b. Any arrangements, understandings, or agreements regarding the
|
| 377 |
+
Licensed Material not stated herein are separate from and
|
| 378 |
+
independent of the terms and conditions of this Public License.
|
| 379 |
+
|
| 380 |
+
|
| 381 |
+
Section 8 -- Interpretation.
|
| 382 |
+
|
| 383 |
+
a. For the avoidance of doubt, this Public License does not, and
|
| 384 |
+
shall not be interpreted to, reduce, limit, restrict, or impose
|
| 385 |
+
conditions on any use of the Licensed Material that could lawfully
|
| 386 |
+
be made without permission under this Public License.
|
| 387 |
+
|
| 388 |
+
b. To the extent possible, if any provision of this Public License is
|
| 389 |
+
deemed unenforceable, it shall be automatically reformed to the
|
| 390 |
+
minimum extent necessary to make it enforceable. If the provision
|
| 391 |
+
cannot be reformed, it shall be severed from this Public License
|
| 392 |
+
without affecting the enforceability of the remaining terms and
|
| 393 |
+
conditions.
|
| 394 |
+
|
| 395 |
+
c. No term or condition of this Public License will be waived and no
|
| 396 |
+
failure to comply consented to unless expressly agreed to by the
|
| 397 |
+
Licensor.
|
| 398 |
+
|
| 399 |
+
d. Nothing in this Public License constitutes or may be interpreted
|
| 400 |
+
as a limitation upon, or waiver of, any privileges and immunities
|
| 401 |
+
that apply to the Licensor or You, including from the legal
|
| 402 |
+
processes of any jurisdiction or authority.
|
| 403 |
+
|
| 404 |
+
=======================================================================
|
| 405 |
+
|
| 406 |
+
Creative Commons is not a party to its public
|
| 407 |
+
licenses. Notwithstanding, Creative Commons may elect to apply one of
|
| 408 |
+
its public licenses to material it publishes and in those instances
|
| 409 |
+
will be considered the “Licensor.” The text of the Creative Commons
|
| 410 |
+
public licenses is dedicated to the public domain under the CC0 Public
|
| 411 |
+
Domain Dedication. Except for the limited purpose of indicating that
|
| 412 |
+
material is shared under a Creative Commons public license or as
|
| 413 |
+
otherwise permitted by the Creative Commons policies published at
|
| 414 |
+
creativecommons.org/policies, Creative Commons does not authorize the
|
| 415 |
+
use of the trademark "Creative Commons" or any other trademark or logo
|
| 416 |
+
of Creative Commons without its prior written consent including,
|
| 417 |
+
without limitation, in connection with any unauthorized modifications
|
| 418 |
+
to any of its public licenses or any other arrangements,
|
| 419 |
+
understandings, or agreements concerning use of licensed material. For
|
| 420 |
+
the avoidance of doubt, this paragraph does not form part of the
|
| 421 |
+
public licenses.
|
| 422 |
+
|
| 423 |
+
Creative Commons may be contacted at creativecommons.org.
|
| 424 |
+
|
README.md
ADDED
|
@@ -0,0 +1,201 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
---
|
| 2 |
+
license: cc-by-nc-4.0
|
| 3 |
+
library_name: transformers
|
| 4 |
+
base_model: m-a-p/MERT-v2-FullSong
|
| 5 |
+
base_model_relation: adapter
|
| 6 |
+
tags:
|
| 7 |
+
- audio
|
| 8 |
+
- music
|
| 9 |
+
- music-transcription
|
| 10 |
+
- midi
|
| 11 |
+
- abc-notation
|
| 12 |
+
- custom_code
|
| 13 |
+
---
|
| 14 |
+
<h1 align="center">🤗 SheetSage2</h1>
|
| 15 |
+
<p align="center"><strong>Music audio to editable scores</strong></p>
|
| 16 |
+
<p align="center">Melody · Chords · Beats · Key · Structure</p>
|
| 17 |
+
|
| 18 |
+
<p align="center">
|
| 19 |
+
<a href="https://map-yue2.github.io/">🎵 YuE2 project</a>
|
| 20 |
+
·
|
| 21 |
+
<a href="#quick-start">🚀 Quick start</a>
|
| 22 |
+
·
|
| 23 |
+
<a href="#benchmarks">📊 Benchmarks</a>
|
| 24 |
+
·
|
| 25 |
+
<a href="#citation">📚 Citation</a>
|
| 26 |
+
</p>
|
| 27 |
+
<p align="center">
|
| 28 |
+
<a href="https://huggingface.co/m-a-p/YuE2-3B"><img alt="🤗 YuE2-3B" src="https://img.shields.io/badge/YuE2--3B-374151?logo=huggingface&logoColor=FFD21E" height="20" /></a>
|
| 29 |
+
|
| 30 |
+
<a href="https://huggingface.co/m-a-p/YuE2-Vae"><img alt="🤗 YuE2-Vae" src="https://img.shields.io/badge/YuE2--Vae-374151?logo=huggingface&logoColor=FFD21E" height="20" /></a>
|
| 31 |
+
|
| 32 |
+
<a href="https://huggingface.co/m-a-p/YuE2-Vae-legacy"><img alt="🤗 YuE2-Vae-legacy" src="https://img.shields.io/badge/YuE2--Vae--legacy-374151?logo=huggingface&logoColor=FFD21E" height="20" /></a>
|
| 33 |
+
|
| 34 |
+
<a href="https://huggingface.co/m-a-p/MERT-v2-30s"><img alt="🤗 MERT-v2-30s" src="https://img.shields.io/badge/MERT--v2--30s-374151?logo=huggingface&logoColor=FFD21E" height="20" /></a>
|
| 35 |
+
|
| 36 |
+
<a href="https://huggingface.co/m-a-p/MERT-v2-FullSong"><img alt="🤗 MERT-v2-FullSong" src="https://img.shields.io/badge/MERT--v2--FullSong-374151?logo=huggingface&logoColor=FFD21E" height="20" /></a>
|
| 37 |
+
|
| 38 |
+
<a href="https://huggingface.co/datasets/m-a-p/WildSongBench"><img alt="🤗 WildSongBench" src="https://img.shields.io/badge/WildSongBench-374151?logo=huggingface&logoColor=FFD21E" height="20" /></a>
|
| 39 |
+
|
| 40 |
+
<a href="https://huggingface.co/m-a-p/SheetSage2"><img alt="SheetSage2" src="https://img.shields.io/badge/SheetSage2-374151?logo=huggingface&logoColor=FFD21E" height="20" /></a>
|
| 41 |
+
</p>
|
| 42 |
+
|
| 43 |
+
**SheetSage2 turns music recordings into lead sheets and timed musical annotations.** Transcribe a complete song, edit its ABC or MIDI, render a piano preview, or extract embeddings and token predictions for your own tools.
|
| 44 |
+
|
| 45 |
+
Built on [MERT-v2-FullSong](https://huggingface.co/m-a-p/MERT-v2-FullSong), with adapters that merge automatically when you load the model.
|
| 46 |
+
|
| 47 |
+

|
| 48 |
+
|
| 49 |
+
<a id="quick-start"></a>
|
| 50 |
+
|
| 51 |
+
## 🚀 Quick start
|
| 52 |
+
|
| 53 |
+
Use Python 3.10 or 3.11 and FFmpeg 6.1 with its shared libraries. Sign in with access to this repository and its MERT-v2 parent:
|
| 54 |
+
|
| 55 |
+
```bash
|
| 56 |
+
python -m pip install huggingface-hub==0.36.0
|
| 57 |
+
huggingface-cli download m-a-p/SheetSage2 --local-dir SheetSage2
|
| 58 |
+
cd SheetSage2
|
| 59 |
+
python -m pip install torch==2.8.0 torchaudio==2.8.0 --index-url https://download.pytorch.org/whl/cu126
|
| 60 |
+
python -m pip install -r requirements.txt
|
| 61 |
+
```
|
| 62 |
+
|
| 63 |
+
```python
|
| 64 |
+
import torch
|
| 65 |
+
from transformers import AutoModel
|
| 66 |
+
|
| 67 |
+
model = AutoModel.from_pretrained(
|
| 68 |
+
"m-a-p/SheetSage2", trust_remote_code=True,
|
| 69 |
+
).eval().to("cuda" if torch.cuda.is_available() else "cpu")
|
| 70 |
+
|
| 71 |
+
result = model.transcribe("song.mp3", output_dir="output")
|
| 72 |
+
```
|
| 73 |
+
|
| 74 |
+
Files are decoded, mixed to mono and resampled automatically. Long songs use overlapping windows.
|
| 75 |
+
|
| 76 |
+
| Output | File |
|
| 77 |
+
|---|---|
|
| 78 |
+
| Editable score | `score.abc` |
|
| 79 |
+
| Melody and chord accompaniment | `transcription.mid` |
|
| 80 |
+
| Separate melodies and chords | `melody_vocal.mid`, `melody_instrumental.mid`, `chords.mid` |
|
| 81 |
+
| Timed annotations | `events.json`, `*.lab` |
|
| 82 |
+
|
| 83 |
+
Omit `output_dir` to keep results in memory. Pass a Tensor or NumPy waveform with its sample rate:
|
| 84 |
+
|
| 85 |
+
```python
|
| 86 |
+
# waveform: [samples] or [channels, samples]
|
| 87 |
+
result = model.transcribe(waveform, sampling_rate=24000)
|
| 88 |
+
abc = result["abc"] # str
|
| 89 |
+
midi = result["midi"] # bytes
|
| 90 |
+
events = result["events"] # list of timed events
|
| 91 |
+
```
|
| 92 |
+
|
| 93 |
+
Paths, encoded audio bytes and binary streams are also accepted. `result["midis"]` contains the separate MIDI parts; `result["labs"]` contains annotation text. Request `export_logits=True`, `export_scores=True` or `export_embeddings=True` for CPU tensors in `result["tensors"]`, grouped by window; add `output_hidden_states=True` for all 24 MERT layers. With `output_dir`, these tensors are saved as safetensors instead. Low-level `forward()` returns logits and optional hidden states; `generate()` returns symbolic tokens.
|
| 94 |
+
|
| 95 |
+
Save a self-contained model for offline use:
|
| 96 |
+
|
| 97 |
+
```python
|
| 98 |
+
model.save_pretrained("sheetsage2-local")
|
| 99 |
+
model = AutoModel.from_pretrained(
|
| 100 |
+
"sheetsage2-local", trust_remote_code=True, local_files_only=True,
|
| 101 |
+
)
|
| 102 |
+
```
|
| 103 |
+
|
| 104 |
+
### Melody-only ABC for covers
|
| 105 |
+
|
| 106 |
+
Set `melody_only=True` to retain both the `Vocal` and `Ins` melodies while omitting chord symbols from the ABC and chord accompaniment from playback/combined MIDI. Raw predicted annotations remain available; the default full transcription is unchanged.
|
| 107 |
+
|
| 108 |
+
```python
|
| 109 |
+
result = model.transcribe("song.mp3", output_dir="cover-score", melody_only=True)
|
| 110 |
+
abc = result["abc"] # Also saved as cover-score/score.abc.
|
| 111 |
+
```
|
| 112 |
+
|
| 113 |
+
```bash
|
| 114 |
+
python infer.py song.mp3 --output cover-score --melody-only
|
| 115 |
+
```
|
| 116 |
+
|
| 117 |
+
Review the transcription, then pass `result["abc"]` or the saved `cover-score/score.abc` to [YuE2](https://huggingface.co/m-a-p/YuE2-3B) as `abc`, with `cot="melody"` and your target `style` and `lyrics`. If the requested ABC cannot be produced, Python raises an error (partial results are available as `error.result`) and the CLI exits with a nonzero status.
|
| 118 |
+
|
| 119 |
+
### Command line and rendering
|
| 120 |
+
|
| 121 |
+
```bash
|
| 122 |
+
python infer.py song.mp3 --output output
|
| 123 |
+
|
| 124 |
+
# Optional piano audio and printable sheet music:
|
| 125 |
+
python setup_render.py
|
| 126 |
+
python infer.py song.mp3 --output output --render-audio --render-score
|
| 127 |
+
|
| 128 |
+
# Render existing results without loading the model:
|
| 129 |
+
python render.py --input output --output rendered --audio --score pdf,svg,png
|
| 130 |
+
```
|
| 131 |
+
|
| 132 |
+
On minimal Linux servers, use `python setup_render.py --with-deps`. Audio follows the original MIDI timing. Scores preserve the vocal and instrumental staves. For separate piano previews, use `--render-parts vocal,instrumental,chords` with `infer.py`, or `--parts vocal,instrumental,chords` with `render.py`.
|
| 133 |
+
|
| 134 |
+
Python also accepts `render_audio=True, render_score="pdf,svg,png"`. In memory mode, `result["rendered"]` contains `audio` (part → WAV bytes) and `score` (format → pages as PDF/PNG bytes or SVG text).
|
| 135 |
+
|
| 136 |
+
<a id="benchmarks"></a>
|
| 137 |
+
|
| 138 |
+
## 📊 Benchmarks
|
| 139 |
+
|
| 140 |
+
Scores (%) on H800 with BF16 inference; higher is better. Use `preset="paper"` or `--preset paper` with the dataset prompts in [scores and settings](benchmark_results.json).
|
| 141 |
+
|
| 142 |
+
| Task | Dataset | Metric | SheetSage1 | madmom | Specialist | SheetSage2 |
|
| 143 |
+
|---|---|---|---:|---:|---:|---:|
|
| 144 |
+
| Beat | GTZAN | F1 | 86.07 | 86.07 | **89.01** [Beat This!][beat] | 86.27 |
|
| 145 |
+
| Beat | osu2017 | F1 | 91.80 | 91.80 | 89.19 [Beat This!][beat] | **93.01** |
|
| 146 |
+
| Downbeat | GTZAN | F1 | 64.65 | 64.65 | 78.28 [Beat This!][beat] | **80.45** |
|
| 147 |
+
| Downbeat | osu2017 | F1 | 83.47 | 83.47 | 85.90 [Beat This!][beat] | **92.90** |
|
| 148 |
+
| Key | GiantSteps | Score | 43.89 | 74.62 | 72.09 [S-KEY][key] | **77.73** |
|
| 149 |
+
| Key | GTZAN | Score | 54.56 | 72.05 | 74.43 [S-KEY][key] | **75.77** |
|
| 150 |
+
| Chord | osu2017 | Maj/min | 79.82 | 77.42 | 84.59 [Jiang et al.][chord] | **90.08** |
|
| 151 |
+
| Chord | Chords1217 | Maj/min | 72.98 | 83.52* | **84.09**† [ChordFormer][chordformer] | 83.81 |
|
| 152 |
+
| Structure | HarmonixSet | Accuracy | — | — | 80.03 [SongFormer][structure] | **80.51** |
|
| 153 |
+
| Structure | HarmonixSet | F1 @ 0.5 s | — | — | **70.63** [SongFormer][structure] | 67.96 |
|
| 154 |
+
| Structure | HarmonixSet | F1 @ 3 s | — | — | 79.50 [SongFormer][structure] | **82.86** |
|
| 155 |
+
| Melody | RWC-Pop | Vocal F1 | 62.71 | — | 62.71 [SheetSage1][sheetsage] | **82.51** |
|
| 156 |
+
| Melody | RWC-Pop | Full F1 | 64.02 | — | 64.02 [SheetSage1][sheetsage] | **75.29** |
|
| 157 |
+
|
| 158 |
+
Beat/downbeat F1 uses a 70 ms matching tolerance and excludes reference and predicted events before 5 s for all models. We evaluate 999 GTZAN tracks for beat, 993 for downbeat, and 142 osu2017 tracks for both. Beat This! uses its original inference settings without DBN postprocessing; scores are averaged over three released training seeds.
|
| 159 |
+
|
| 160 |
+
Melody F1 uses pitch classes. SheetSage1 beat/downbeat results use madmom. *The madmom chord model includes Chords1217 in training. †ChordFormer uses five-fold cross-validation; SheetSage2 uses one checkpoint across all tracks.
|
| 161 |
+
|
| 162 |
+
### Benchmark updates
|
| 163 |
+
|
| 164 |
+
- **2026-09-14:** Updated beat/downbeat scores with a shared 5-second exclusion protocol. Beat This! was rerun under its original settings, averaging scores across three released seeds. Key metric labels are now consistent with the report (Score).
|
| 165 |
+
|
| 166 |
+
[beat]: https://doi.org/10.5281/zenodo.14877491
|
| 167 |
+
[key]: https://doi.org/10.1109/ICASSP49660.2025.10890222
|
| 168 |
+
[chord]: https://archives.ismir.net/ismir2019/paper/000078.pdf
|
| 169 |
+
[chordformer]: https://arxiv.org/abs/2502.11840
|
| 170 |
+
[structure]: https://arxiv.org/abs/2510.02797
|
| 171 |
+
[sheetsage]: https://arxiv.org/abs/2212.01884
|
| 172 |
+
|
| 173 |
+
<a id="citation"></a>
|
| 174 |
+
|
| 175 |
+
## 📚 Citation
|
| 176 |
+
|
| 177 |
+
**Technical report coming soon.** For now, please cite [YuE](https://arxiv.org/abs/2503.08638) and [MERT](https://proceedings.iclr.cc/paper_files/paper/2024/hash/33dffa2e3d2ab74a783d1a8c292f66d9-Abstract-Conference.html) when using SheetSage2 in your research.
|
| 178 |
+
|
| 179 |
+
```bibtex
|
| 180 |
+
@article{yuan2025yue,
|
| 181 |
+
title = {{YuE}: Scaling Open Foundation Models for Long-Form Music Generation},
|
| 182 |
+
author = {Yuan, Ruibin and Lin, Hanfeng and Guo, Shuyue and Zhang, Ge and Pan, Jiahao and Zang, Yongyi and Liu, Haohe and Liang, Yiming and Ma, Wenye and Du, Xingjian and Du, Xinrun and Ye, Zhen and Zheng, Tianyu and Jiang, Zhengxuan and Ma, Yinghao and Liu, Minghao and Tian, Zeyue and Zhou, Ziya and Xue, Liumeng and Qu, Xingwei and Li, Yizhi and Wu, Shangda and Shen, Tianhao and Ma, Ziyang and Zhan, Jun and Wang, Chunhui and Wang, Yatian and Chi, Xiaowei and Zhang, Xinyue and Yang, Zhenzhu and Wang, Xiangzhou and Liu, Shansong and Mei, Lingrui and Li, Peng and Wang, Junjie and Yu, Jianwei and Pang, Guojian and Li, Xu and Wang, Zihao and Zhou, Xiaohuan and Yu, Lijun and Benetos, Emmanouil and Chen, Yong and Lin, Chenghua and Chen, Xie and Xia, Gus and Zhang, Zhaoxiang and Zhang, Chao and Chen, Wenhu and Zhou, Xinyu and Qiu, Xipeng and Dannenberg, Roger and Liu, Jiaheng and Yang, Jian and Huang, Wenhao and Xue, Wei and Tan, Xu and Guo, Yike},
|
| 183 |
+
journal = {arXiv preprint arXiv:2503.08638},
|
| 184 |
+
year = {2025},
|
| 185 |
+
eprint = {2503.08638},
|
| 186 |
+
archivePrefix = {arXiv},
|
| 187 |
+
url = {https://arxiv.org/abs/2503.08638}
|
| 188 |
+
}
|
| 189 |
+
|
| 190 |
+
@inproceedings{li2024mert,
|
| 191 |
+
title = {MERT: Acoustic Music Understanding Model with Large-Scale Self-supervised Training},
|
| 192 |
+
author = {Li, Yizhi and Yuan, Ruibin and Zhang, Ge and Ma, Yinghao and Chen, Xingran and Yin, Hanzhi and Xiao, Chenghao and Lin, Chenghua and Ragni, Anton and Benetos, Emmanouil and Gyenge, Norbert and Dannenberg, Roger and Liu, Ruibo and Chen, Wenhu and Xia, Gus and Shi, Yemin and Huang, Wenhao and Wang, Zili and Guo, Yike and Fu, Jie},
|
| 193 |
+
booktitle = {International Conference on Learning Representations},
|
| 194 |
+
year = {2024},
|
| 195 |
+
url = {https://proceedings.iclr.cc/paper_files/paper/2024/hash/33dffa2e3d2ab74a783d1a8c292f66d9-Abstract-Conference.html}
|
| 196 |
+
}
|
| 197 |
+
```
|
| 198 |
+
|
| 199 |
+
Weights: [CC BY-NC 4.0](LICENSE). [Third-party notices](THIRD_PARTY_NOTICES.md).
|
| 200 |
+
|
| 201 |
+
**YuE2 family:** [Song generation](https://huggingface.co/m-a-p/YuE2-3B) · [Music representations](https://huggingface.co/m-a-p/MERT-v2-FullSong) · [Audio decoder](https://huggingface.co/m-a-p/YuE2-Vae).
|
THIRD_PARTY_NOTICES.md
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Third-party notices
|
| 2 |
+
|
| 3 |
+
- MERT-v2-FullSong provides the public music encoder. Its weight terms remain CC BY-NC 4.0.
|
| 4 |
+
- The BART decoder uses Hugging Face Transformers (Apache 2.0); PyTorch uses its BSD-style license.
|
| 5 |
+
- abcjs 6.6.3 is distributed under MIT; its license accompanies the rendering assets.
|
| 6 |
+
- Playwright is distributed under Apache 2.0. Chromium is installed separately with its accompanying notices.
|
| 7 |
+
- FluidR3 piano samples by Frank Wen use CC BY 3.0 US. Attribution and source information accompany the samples.
|
| 8 |
+
- Other installed dependencies retain their respective licenses.
|
__init__.py
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
"""SheetSage2 standalone inference."""
|
assets/architecture.png
ADDED
|
Git LFS Details
|
audio_sheetsage2.py
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""File decoding and waveform preparation for music transcription."""
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
import io
|
| 4 |
+
import math
|
| 5 |
+
import numbers
|
| 6 |
+
import shutil
|
| 7 |
+
import subprocess
|
| 8 |
+
|
| 9 |
+
import numpy as np
|
| 10 |
+
import torch
|
| 11 |
+
|
| 12 |
+
SAMPLE_RATE = 24000
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def _encoded_bytes(audio):
|
| 16 |
+
if isinstance(audio, (bytes, bytearray, memoryview)):
|
| 17 |
+
return bytes(audio)
|
| 18 |
+
if callable(getattr(audio, "read", None)):
|
| 19 |
+
position = None
|
| 20 |
+
try:
|
| 21 |
+
position = audio.tell()
|
| 22 |
+
audio.seek(0)
|
| 23 |
+
except (AttributeError, OSError, ValueError, io.UnsupportedOperation):
|
| 24 |
+
position = None
|
| 25 |
+
try:
|
| 26 |
+
value = audio.read()
|
| 27 |
+
finally:
|
| 28 |
+
if position is not None:
|
| 29 |
+
audio.seek(position)
|
| 30 |
+
if not isinstance(value, (bytes, bytearray)):
|
| 31 |
+
raise ValueError("Audio streams must return encoded audio bytes")
|
| 32 |
+
return bytes(value)
|
| 33 |
+
return None
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def load_audio(audio, *, sampling_rate=None, max_seconds=None, preset="default"):
|
| 37 |
+
if max_seconds is not None and (not math.isfinite(max_seconds) or max_seconds <= 0):
|
| 38 |
+
raise ValueError("max_seconds must be finite and positive")
|
| 39 |
+
encoded = _encoded_bytes(audio)
|
| 40 |
+
if encoded is not None and not encoded:
|
| 41 |
+
raise ValueError("Encoded audio must not be empty")
|
| 42 |
+
if isinstance(audio, (str, Path)) or encoded is not None:
|
| 43 |
+
if preset == "paper":
|
| 44 |
+
import torchaudio
|
| 45 |
+
source = io.BytesIO(encoded) if encoded is not None else str(audio)
|
| 46 |
+
info = torchaudio.info(source, backend="ffmpeg")
|
| 47 |
+
if info.num_frames <= 0:
|
| 48 |
+
raise ValueError("Cannot determine source audio length")
|
| 49 |
+
if encoded is not None:
|
| 50 |
+
source.seek(0)
|
| 51 |
+
waveform, rate = torchaudio.load(source, frame_offset=0, num_frames=int(info.num_frames),
|
| 52 |
+
backend="ffmpeg", channels_first=True)
|
| 53 |
+
waveform = waveform.mean(dim=0)
|
| 54 |
+
if rate != SAMPLE_RATE:
|
| 55 |
+
waveform = torchaudio.functional.resample(waveform, rate, SAMPLE_RATE)
|
| 56 |
+
else:
|
| 57 |
+
if not shutil.which("ffmpeg"):
|
| 58 |
+
raise RuntimeError("FFmpeg is required to read audio files; install it and add it to PATH")
|
| 59 |
+
source = "pipe:0" if encoded is not None else str(Path(audio).resolve())
|
| 60 |
+
command = ["ffmpeg", "-v", "error", "-nostdin", "-i", source, "-vn"]
|
| 61 |
+
if max_seconds is not None:
|
| 62 |
+
command += ["-t", str(float(max_seconds))]
|
| 63 |
+
command += ["-ac", "1", "-ar", str(SAMPLE_RATE), "-f", "f32le", "pipe:1"]
|
| 64 |
+
result = subprocess.run(command, input=encoded, capture_output=True, timeout=600, check=False)
|
| 65 |
+
if result.returncode:
|
| 66 |
+
raise ValueError("Cannot decode audio: " + result.stderr.decode(errors="replace")[-1200:])
|
| 67 |
+
waveform = torch.from_numpy(np.frombuffer(result.stdout, dtype="<f4").copy())
|
| 68 |
+
else:
|
| 69 |
+
if not isinstance(sampling_rate, numbers.Real) or not math.isfinite(sampling_rate) or sampling_rate <= 0 or sampling_rate != int(sampling_rate):
|
| 70 |
+
raise ValueError("Provide sampling_rate for an array or tensor waveform")
|
| 71 |
+
waveform = torch.as_tensor(audio, dtype=torch.float32, device="cpu")
|
| 72 |
+
if waveform.ndim == 2:
|
| 73 |
+
if waveform.shape[0] > 32:
|
| 74 |
+
raise ValueError("Multichannel audio must have shape [channels, samples]")
|
| 75 |
+
waveform = waveform.mean(dim=0)
|
| 76 |
+
if waveform.ndim != 1:
|
| 77 |
+
raise ValueError("Audio must have shape [samples] or [channels, samples]")
|
| 78 |
+
if sampling_rate != SAMPLE_RATE:
|
| 79 |
+
import torchaudio
|
| 80 |
+
waveform = torchaudio.functional.resample(waveform, int(sampling_rate), SAMPLE_RATE)
|
| 81 |
+
waveform = waveform.float().contiguous()
|
| 82 |
+
if max_seconds is not None:
|
| 83 |
+
waveform = waveform[:round(max_seconds * SAMPLE_RATE)]
|
| 84 |
+
if waveform.numel() < 1025 or not torch.isfinite(waveform).all():
|
| 85 |
+
raise ValueError("Audio must contain at least 1025 finite samples at 24 kHz")
|
| 86 |
+
return waveform
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def slice_audio(audio, start, seconds):
|
| 90 |
+
offset, count = round(start * SAMPLE_RATE), round(seconds * SAMPLE_RATE)
|
| 91 |
+
result = audio[offset:offset + count]
|
| 92 |
+
return torch.nn.functional.pad(result, (0, count - len(result)))
|
benchmark_results.json
ADDED
|
@@ -0,0 +1,379 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"scores": [
|
| 3 |
+
{
|
| 4 |
+
"task": "beat",
|
| 5 |
+
"dataset": "GTZAN",
|
| 6 |
+
"metric": "F1",
|
| 7 |
+
"score": 86.27,
|
| 8 |
+
"settings": {
|
| 9 |
+
"preset": "paper",
|
| 10 |
+
"sample_rate": 24000,
|
| 11 |
+
"window_seconds": 300,
|
| 12 |
+
"overlap_seconds": 100,
|
| 13 |
+
"lookahead_seconds": 0,
|
| 14 |
+
"max_length": 5120,
|
| 15 |
+
"decoding": "greedy",
|
| 16 |
+
"attention": "sdpa",
|
| 17 |
+
"weight_dtype": "float32",
|
| 18 |
+
"autocast_dtype": "bfloat16",
|
| 19 |
+
"batch_size": 1,
|
| 20 |
+
"prompts": [
|
| 21 |
+
"timestamp",
|
| 22 |
+
"downbeat_meter",
|
| 23 |
+
"key"
|
| 24 |
+
],
|
| 25 |
+
"device": "NVIDIA H800",
|
| 26 |
+
"cuda_version": "12.6"
|
| 27 |
+
},
|
| 28 |
+
"evaluation": {
|
| 29 |
+
"num_tracks": 999,
|
| 30 |
+
"matching_tolerance_seconds": 0.07,
|
| 31 |
+
"min_event_time_seconds": 5.0,
|
| 32 |
+
"trim_reference_and_predictions": true,
|
| 33 |
+
"aggregation": "macro_average_per_track_f1"
|
| 34 |
+
}
|
| 35 |
+
},
|
| 36 |
+
{
|
| 37 |
+
"task": "beat",
|
| 38 |
+
"dataset": "osu2017",
|
| 39 |
+
"metric": "F1",
|
| 40 |
+
"score": 93.01,
|
| 41 |
+
"settings": {
|
| 42 |
+
"preset": "paper",
|
| 43 |
+
"sample_rate": 24000,
|
| 44 |
+
"window_seconds": 300,
|
| 45 |
+
"overlap_seconds": 100,
|
| 46 |
+
"lookahead_seconds": 0,
|
| 47 |
+
"max_length": 5120,
|
| 48 |
+
"decoding": "greedy",
|
| 49 |
+
"attention": "sdpa",
|
| 50 |
+
"weight_dtype": "float32",
|
| 51 |
+
"autocast_dtype": "bfloat16",
|
| 52 |
+
"batch_size": 1,
|
| 53 |
+
"prompts": [
|
| 54 |
+
"timestamp",
|
| 55 |
+
"downbeat_meter",
|
| 56 |
+
"key",
|
| 57 |
+
"chord_full"
|
| 58 |
+
],
|
| 59 |
+
"device": "NVIDIA H800",
|
| 60 |
+
"cuda_version": "12.6"
|
| 61 |
+
},
|
| 62 |
+
"evaluation": {
|
| 63 |
+
"num_tracks": 142,
|
| 64 |
+
"matching_tolerance_seconds": 0.07,
|
| 65 |
+
"min_event_time_seconds": 5.0,
|
| 66 |
+
"trim_reference_and_predictions": true,
|
| 67 |
+
"aggregation": "macro_average_per_track_f1"
|
| 68 |
+
}
|
| 69 |
+
},
|
| 70 |
+
{
|
| 71 |
+
"task": "downbeat",
|
| 72 |
+
"dataset": "GTZAN",
|
| 73 |
+
"metric": "F1",
|
| 74 |
+
"score": 80.45,
|
| 75 |
+
"settings": {
|
| 76 |
+
"preset": "paper",
|
| 77 |
+
"sample_rate": 24000,
|
| 78 |
+
"window_seconds": 300,
|
| 79 |
+
"overlap_seconds": 100,
|
| 80 |
+
"lookahead_seconds": 0,
|
| 81 |
+
"max_length": 5120,
|
| 82 |
+
"decoding": "greedy",
|
| 83 |
+
"attention": "sdpa",
|
| 84 |
+
"weight_dtype": "float32",
|
| 85 |
+
"autocast_dtype": "bfloat16",
|
| 86 |
+
"batch_size": 1,
|
| 87 |
+
"prompts": [
|
| 88 |
+
"timestamp",
|
| 89 |
+
"downbeat_meter",
|
| 90 |
+
"key"
|
| 91 |
+
],
|
| 92 |
+
"device": "NVIDIA H800",
|
| 93 |
+
"cuda_version": "12.6"
|
| 94 |
+
},
|
| 95 |
+
"evaluation": {
|
| 96 |
+
"num_tracks": 993,
|
| 97 |
+
"matching_tolerance_seconds": 0.07,
|
| 98 |
+
"min_event_time_seconds": 5.0,
|
| 99 |
+
"trim_reference_and_predictions": true,
|
| 100 |
+
"aggregation": "macro_average_per_track_f1"
|
| 101 |
+
}
|
| 102 |
+
},
|
| 103 |
+
{
|
| 104 |
+
"task": "downbeat",
|
| 105 |
+
"dataset": "osu2017",
|
| 106 |
+
"metric": "F1",
|
| 107 |
+
"score": 92.9,
|
| 108 |
+
"settings": {
|
| 109 |
+
"preset": "paper",
|
| 110 |
+
"sample_rate": 24000,
|
| 111 |
+
"window_seconds": 300,
|
| 112 |
+
"overlap_seconds": 100,
|
| 113 |
+
"lookahead_seconds": 0,
|
| 114 |
+
"max_length": 5120,
|
| 115 |
+
"decoding": "greedy",
|
| 116 |
+
"attention": "sdpa",
|
| 117 |
+
"weight_dtype": "float32",
|
| 118 |
+
"autocast_dtype": "bfloat16",
|
| 119 |
+
"batch_size": 1,
|
| 120 |
+
"prompts": [
|
| 121 |
+
"timestamp",
|
| 122 |
+
"downbeat_meter",
|
| 123 |
+
"key",
|
| 124 |
+
"chord_full"
|
| 125 |
+
],
|
| 126 |
+
"device": "NVIDIA H800",
|
| 127 |
+
"cuda_version": "12.6"
|
| 128 |
+
},
|
| 129 |
+
"evaluation": {
|
| 130 |
+
"num_tracks": 142,
|
| 131 |
+
"matching_tolerance_seconds": 0.07,
|
| 132 |
+
"min_event_time_seconds": 5.0,
|
| 133 |
+
"trim_reference_and_predictions": true,
|
| 134 |
+
"aggregation": "macro_average_per_track_f1"
|
| 135 |
+
}
|
| 136 |
+
},
|
| 137 |
+
{
|
| 138 |
+
"task": "key",
|
| 139 |
+
"dataset": "GiantSteps",
|
| 140 |
+
"metric": "weighted_accuracy",
|
| 141 |
+
"score": 77.73,
|
| 142 |
+
"settings": {
|
| 143 |
+
"preset": "paper",
|
| 144 |
+
"sample_rate": 24000,
|
| 145 |
+
"window_seconds": 300,
|
| 146 |
+
"overlap_seconds": 100,
|
| 147 |
+
"lookahead_seconds": 0,
|
| 148 |
+
"max_length": 5120,
|
| 149 |
+
"decoding": "greedy",
|
| 150 |
+
"attention": "sdpa",
|
| 151 |
+
"weight_dtype": "float32",
|
| 152 |
+
"autocast_dtype": "bfloat16",
|
| 153 |
+
"batch_size": 1,
|
| 154 |
+
"prompts": [
|
| 155 |
+
"timestamp",
|
| 156 |
+
"downbeat_meter",
|
| 157 |
+
"key"
|
| 158 |
+
],
|
| 159 |
+
"device": "NVIDIA H800",
|
| 160 |
+
"cuda_version": "12.6"
|
| 161 |
+
}
|
| 162 |
+
},
|
| 163 |
+
{
|
| 164 |
+
"task": "key",
|
| 165 |
+
"dataset": "GTZAN",
|
| 166 |
+
"metric": "weighted_accuracy",
|
| 167 |
+
"score": 75.77,
|
| 168 |
+
"settings": {
|
| 169 |
+
"preset": "paper",
|
| 170 |
+
"sample_rate": 24000,
|
| 171 |
+
"window_seconds": 300,
|
| 172 |
+
"overlap_seconds": 100,
|
| 173 |
+
"lookahead_seconds": 0,
|
| 174 |
+
"max_length": 5120,
|
| 175 |
+
"decoding": "greedy",
|
| 176 |
+
"attention": "sdpa",
|
| 177 |
+
"weight_dtype": "float32",
|
| 178 |
+
"autocast_dtype": "bfloat16",
|
| 179 |
+
"batch_size": 1,
|
| 180 |
+
"prompts": [
|
| 181 |
+
"timestamp",
|
| 182 |
+
"downbeat_meter",
|
| 183 |
+
"key"
|
| 184 |
+
],
|
| 185 |
+
"device": "NVIDIA H800",
|
| 186 |
+
"cuda_version": "12.6"
|
| 187 |
+
}
|
| 188 |
+
},
|
| 189 |
+
{
|
| 190 |
+
"task": "chord",
|
| 191 |
+
"dataset": "osu2017",
|
| 192 |
+
"metric": "majmin",
|
| 193 |
+
"score": 90.08,
|
| 194 |
+
"settings": {
|
| 195 |
+
"preset": "paper",
|
| 196 |
+
"sample_rate": 24000,
|
| 197 |
+
"window_seconds": 300,
|
| 198 |
+
"overlap_seconds": 100,
|
| 199 |
+
"lookahead_seconds": 0,
|
| 200 |
+
"max_length": 5120,
|
| 201 |
+
"decoding": "greedy",
|
| 202 |
+
"attention": "sdpa",
|
| 203 |
+
"weight_dtype": "float32",
|
| 204 |
+
"autocast_dtype": "bfloat16",
|
| 205 |
+
"batch_size": 1,
|
| 206 |
+
"prompts": [
|
| 207 |
+
"timestamp",
|
| 208 |
+
"downbeat_meter",
|
| 209 |
+
"key",
|
| 210 |
+
"chord_full"
|
| 211 |
+
],
|
| 212 |
+
"device": "NVIDIA H800",
|
| 213 |
+
"cuda_version": "12.6"
|
| 214 |
+
}
|
| 215 |
+
},
|
| 216 |
+
{
|
| 217 |
+
"task": "chord",
|
| 218 |
+
"dataset": "Chords1217",
|
| 219 |
+
"metric": "majmin",
|
| 220 |
+
"score": 83.81,
|
| 221 |
+
"settings": {
|
| 222 |
+
"preset": "paper",
|
| 223 |
+
"sample_rate": 24000,
|
| 224 |
+
"window_seconds": 300,
|
| 225 |
+
"overlap_seconds": 100,
|
| 226 |
+
"lookahead_seconds": 0,
|
| 227 |
+
"max_length": 5120,
|
| 228 |
+
"decoding": "greedy",
|
| 229 |
+
"attention": "sdpa",
|
| 230 |
+
"weight_dtype": "float32",
|
| 231 |
+
"autocast_dtype": "bfloat16",
|
| 232 |
+
"batch_size": 1,
|
| 233 |
+
"prompts": [
|
| 234 |
+
"timestamp",
|
| 235 |
+
"downbeat_meter",
|
| 236 |
+
"chord_full"
|
| 237 |
+
],
|
| 238 |
+
"device": "NVIDIA H800",
|
| 239 |
+
"cuda_version": "12.6"
|
| 240 |
+
}
|
| 241 |
+
},
|
| 242 |
+
{
|
| 243 |
+
"task": "structure",
|
| 244 |
+
"dataset": "HarmonixSet",
|
| 245 |
+
"metric": "accuracy",
|
| 246 |
+
"score": 80.51,
|
| 247 |
+
"settings": {
|
| 248 |
+
"preset": "paper",
|
| 249 |
+
"sample_rate": 24000,
|
| 250 |
+
"window_seconds": 300,
|
| 251 |
+
"overlap_seconds": 100,
|
| 252 |
+
"lookahead_seconds": 0,
|
| 253 |
+
"max_length": 5120,
|
| 254 |
+
"decoding": "greedy",
|
| 255 |
+
"attention": "sdpa",
|
| 256 |
+
"weight_dtype": "float32",
|
| 257 |
+
"autocast_dtype": "bfloat16",
|
| 258 |
+
"batch_size": 1,
|
| 259 |
+
"prompts": [
|
| 260 |
+
"timestamp",
|
| 261 |
+
"downbeat_meter",
|
| 262 |
+
"structure"
|
| 263 |
+
],
|
| 264 |
+
"device": "NVIDIA H800",
|
| 265 |
+
"cuda_version": "12.6"
|
| 266 |
+
}
|
| 267 |
+
},
|
| 268 |
+
{
|
| 269 |
+
"task": "structure",
|
| 270 |
+
"dataset": "HarmonixSet",
|
| 271 |
+
"metric": "F1@0.5s",
|
| 272 |
+
"score": 67.96,
|
| 273 |
+
"settings": {
|
| 274 |
+
"preset": "paper",
|
| 275 |
+
"sample_rate": 24000,
|
| 276 |
+
"window_seconds": 300,
|
| 277 |
+
"overlap_seconds": 100,
|
| 278 |
+
"lookahead_seconds": 0,
|
| 279 |
+
"max_length": 5120,
|
| 280 |
+
"decoding": "greedy",
|
| 281 |
+
"attention": "sdpa",
|
| 282 |
+
"weight_dtype": "float32",
|
| 283 |
+
"autocast_dtype": "bfloat16",
|
| 284 |
+
"batch_size": 1,
|
| 285 |
+
"prompts": [
|
| 286 |
+
"timestamp",
|
| 287 |
+
"downbeat_meter",
|
| 288 |
+
"structure"
|
| 289 |
+
],
|
| 290 |
+
"device": "NVIDIA H800",
|
| 291 |
+
"cuda_version": "12.6"
|
| 292 |
+
}
|
| 293 |
+
},
|
| 294 |
+
{
|
| 295 |
+
"task": "structure",
|
| 296 |
+
"dataset": "HarmonixSet",
|
| 297 |
+
"metric": "F1@3s",
|
| 298 |
+
"score": 82.86,
|
| 299 |
+
"settings": {
|
| 300 |
+
"preset": "paper",
|
| 301 |
+
"sample_rate": 24000,
|
| 302 |
+
"window_seconds": 300,
|
| 303 |
+
"overlap_seconds": 100,
|
| 304 |
+
"lookahead_seconds": 0,
|
| 305 |
+
"max_length": 5120,
|
| 306 |
+
"decoding": "greedy",
|
| 307 |
+
"attention": "sdpa",
|
| 308 |
+
"weight_dtype": "float32",
|
| 309 |
+
"autocast_dtype": "bfloat16",
|
| 310 |
+
"batch_size": 1,
|
| 311 |
+
"prompts": [
|
| 312 |
+
"timestamp",
|
| 313 |
+
"downbeat_meter",
|
| 314 |
+
"structure"
|
| 315 |
+
],
|
| 316 |
+
"device": "NVIDIA H800",
|
| 317 |
+
"cuda_version": "12.6"
|
| 318 |
+
}
|
| 319 |
+
},
|
| 320 |
+
{
|
| 321 |
+
"task": "melody",
|
| 322 |
+
"dataset": "RWC-Pop",
|
| 323 |
+
"metric": "vocal_pitch_class_F1",
|
| 324 |
+
"score": 82.51,
|
| 325 |
+
"settings": {
|
| 326 |
+
"preset": "paper",
|
| 327 |
+
"sample_rate": 24000,
|
| 328 |
+
"window_seconds": 300,
|
| 329 |
+
"overlap_seconds": 100,
|
| 330 |
+
"lookahead_seconds": 0,
|
| 331 |
+
"max_length": 5120,
|
| 332 |
+
"decoding": "greedy",
|
| 333 |
+
"attention": "sdpa",
|
| 334 |
+
"weight_dtype": "float32",
|
| 335 |
+
"autocast_dtype": "bfloat16",
|
| 336 |
+
"batch_size": 1,
|
| 337 |
+
"prompts": [
|
| 338 |
+
"timestamp",
|
| 339 |
+
"downbeat_meter",
|
| 340 |
+
"structure",
|
| 341 |
+
"key",
|
| 342 |
+
"chord_full",
|
| 343 |
+
"melody_full"
|
| 344 |
+
],
|
| 345 |
+
"device": "NVIDIA H800",
|
| 346 |
+
"cuda_version": "12.6"
|
| 347 |
+
}
|
| 348 |
+
},
|
| 349 |
+
{
|
| 350 |
+
"task": "melody",
|
| 351 |
+
"dataset": "RWC-Pop",
|
| 352 |
+
"metric": "full_pitch_class_F1",
|
| 353 |
+
"score": 75.29,
|
| 354 |
+
"settings": {
|
| 355 |
+
"preset": "paper",
|
| 356 |
+
"sample_rate": 24000,
|
| 357 |
+
"window_seconds": 300,
|
| 358 |
+
"overlap_seconds": 100,
|
| 359 |
+
"lookahead_seconds": 0,
|
| 360 |
+
"max_length": 5120,
|
| 361 |
+
"decoding": "greedy",
|
| 362 |
+
"attention": "sdpa",
|
| 363 |
+
"weight_dtype": "float32",
|
| 364 |
+
"autocast_dtype": "bfloat16",
|
| 365 |
+
"batch_size": 1,
|
| 366 |
+
"prompts": [
|
| 367 |
+
"timestamp",
|
| 368 |
+
"downbeat_meter",
|
| 369 |
+
"structure",
|
| 370 |
+
"key",
|
| 371 |
+
"chord_full",
|
| 372 |
+
"melody_full"
|
| 373 |
+
],
|
| 374 |
+
"device": "NVIDIA H800",
|
| 375 |
+
"cuda_version": "12.6"
|
| 376 |
+
}
|
| 377 |
+
}
|
| 378 |
+
]
|
| 379 |
+
}
|
config.json
ADDED
|
@@ -0,0 +1,75 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"architectures": [
|
| 3 |
+
"SheetSage2Model"
|
| 4 |
+
],
|
| 5 |
+
"auto_map": {
|
| 6 |
+
"AutoConfig": "configuration_sheetsage2.SheetSage2Config",
|
| 7 |
+
"AutoModel": "modeling_sheetsage2.SheetSage2Model",
|
| 8 |
+
"AutoModelForSeq2SeqLM": "modeling_sheetsage2.SheetSage2Model",
|
| 9 |
+
"AutoProcessor": "processing_sheetsage2.SheetSage2Processor"
|
| 10 |
+
},
|
| 11 |
+
"backbone_config": {
|
| 12 |
+
"architectures": [
|
| 13 |
+
"MERT2Model"
|
| 14 |
+
],
|
| 15 |
+
"context_seconds": 360.0,
|
| 16 |
+
"conv_depthwise_kernel_size": 31,
|
| 17 |
+
"frame_rate": 25.0,
|
| 18 |
+
"hidden_size": 1024,
|
| 19 |
+
"hop_length": 240,
|
| 20 |
+
"initializer_range": 0.02,
|
| 21 |
+
"inputs_to_logits_ratio": 960,
|
| 22 |
+
"intermediate_size": 4096,
|
| 23 |
+
"layer_norm_eps": 1e-05,
|
| 24 |
+
"minimum_input_samples": 1025,
|
| 25 |
+
"model_type": "mert2",
|
| 26 |
+
"n_fft": 2048,
|
| 27 |
+
"num_attention_heads": 16,
|
| 28 |
+
"num_hidden_layers": 24,
|
| 29 |
+
"num_mel_bins": 128,
|
| 30 |
+
"rotary_embedding_base": 10000,
|
| 31 |
+
"sampling_rate": 24000,
|
| 32 |
+
"subsampling_channels": [
|
| 33 |
+
128,
|
| 34 |
+
512,
|
| 35 |
+
1024
|
| 36 |
+
],
|
| 37 |
+
"subsampling_depths": [
|
| 38 |
+
3,
|
| 39 |
+
4,
|
| 40 |
+
5
|
| 41 |
+
],
|
| 42 |
+
"subsampling_layer_norm_eps": 1e-06,
|
| 43 |
+
"torch_dtype": "float32",
|
| 44 |
+
"transformers_version": "4.53.2",
|
| 45 |
+
"variant": "fs",
|
| 46 |
+
"win_length": 2048
|
| 47 |
+
},
|
| 48 |
+
"base_model_name_or_path": "m-a-p/MERT-v2-FullSong",
|
| 49 |
+
"base_model_revision": "d8ba1c745e733b3908ce6ad16ebeb17ac7600a42",
|
| 50 |
+
"base_model_sha256": "e6dd2ab187d6dd62b6521cd7d8f932e237acf0c5757745a7232082e28391350d",
|
| 51 |
+
"bos_token_id": 1,
|
| 52 |
+
"decoder_dropout": 0.1,
|
| 53 |
+
"decoder_layers": 6,
|
| 54 |
+
"decoder_start_token_id": 1,
|
| 55 |
+
"encoder_attn_implementation": "sdpa",
|
| 56 |
+
"eos_token_id": 2,
|
| 57 |
+
"hidden_size": 512,
|
| 58 |
+
"input_audio_length": 300.0,
|
| 59 |
+
"intermediate_size": 2048,
|
| 60 |
+
"is_encoder_decoder": true,
|
| 61 |
+
"lora_alpha": 128.0,
|
| 62 |
+
"lora_rank": 64,
|
| 63 |
+
"max_output_seq_len": 5120,
|
| 64 |
+
"model_type": "sheetsage2",
|
| 65 |
+
"num_attention_heads": 8,
|
| 66 |
+
"pad_token_id": 0,
|
| 67 |
+
"sampling_rate": 24000,
|
| 68 |
+
"time_hz": 100,
|
| 69 |
+
"tokenizer_fingerprint": "5ba3325af0344c7f",
|
| 70 |
+
"tokenizer_schema_version": "v1",
|
| 71 |
+
"transformers_version": "4.45.2",
|
| 72 |
+
"use_cache": true,
|
| 73 |
+
"vocab_size": 31678,
|
| 74 |
+
"weights_format": "adapter"
|
| 75 |
+
}
|
configuration_mert2.py
ADDED
|
@@ -0,0 +1,76 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Configuration for MERT2 music representation models."""
|
| 2 |
+
|
| 3 |
+
from transformers import PretrainedConfig
|
| 4 |
+
|
| 5 |
+
|
| 6 |
+
class MERT2Config(PretrainedConfig):
|
| 7 |
+
"""Architecture shared by the 30-second and full-song MERT2 encoders."""
|
| 8 |
+
|
| 9 |
+
model_type = "mert2"
|
| 10 |
+
|
| 11 |
+
def __init__(
|
| 12 |
+
self,
|
| 13 |
+
hidden_size=1024,
|
| 14 |
+
intermediate_size=4096,
|
| 15 |
+
num_hidden_layers=24,
|
| 16 |
+
num_attention_heads=16,
|
| 17 |
+
num_mel_bins=128,
|
| 18 |
+
sampling_rate=24000,
|
| 19 |
+
n_fft=2048,
|
| 20 |
+
win_length=2048,
|
| 21 |
+
hop_length=240,
|
| 22 |
+
subsampling_channels=None,
|
| 23 |
+
subsampling_depths=None,
|
| 24 |
+
conv_depthwise_kernel_size=31,
|
| 25 |
+
rotary_embedding_base=10000,
|
| 26 |
+
layer_norm_eps=1e-5,
|
| 27 |
+
subsampling_layer_norm_eps=1e-6,
|
| 28 |
+
initializer_range=0.02,
|
| 29 |
+
variant="30s",
|
| 30 |
+
context_seconds=None,
|
| 31 |
+
**kwargs,
|
| 32 |
+
):
|
| 33 |
+
super().__init__(**kwargs)
|
| 34 |
+
self.hidden_size = int(hidden_size)
|
| 35 |
+
self.intermediate_size = int(intermediate_size)
|
| 36 |
+
self.num_hidden_layers = int(num_hidden_layers)
|
| 37 |
+
self.num_attention_heads = int(num_attention_heads)
|
| 38 |
+
self.num_mel_bins = int(num_mel_bins)
|
| 39 |
+
self.sampling_rate = int(sampling_rate)
|
| 40 |
+
self.n_fft = int(n_fft)
|
| 41 |
+
self.win_length = int(win_length)
|
| 42 |
+
self.hop_length = int(hop_length)
|
| 43 |
+
self.subsampling_channels = list(subsampling_channels or [num_mel_bins, 512, hidden_size])
|
| 44 |
+
self.subsampling_depths = list(subsampling_depths or [3, 4, 5])
|
| 45 |
+
self.conv_depthwise_kernel_size = int(conv_depthwise_kernel_size)
|
| 46 |
+
self.rotary_embedding_base = int(rotary_embedding_base)
|
| 47 |
+
self.layer_norm_eps = float(layer_norm_eps)
|
| 48 |
+
self.subsampling_layer_norm_eps = float(subsampling_layer_norm_eps)
|
| 49 |
+
self.initializer_range = float(initializer_range)
|
| 50 |
+
self.variant = str(variant)
|
| 51 |
+
self.context_seconds = float(context_seconds if context_seconds is not None else (360 if variant == "fs" else 30))
|
| 52 |
+
self.inputs_to_logits_ratio = self.hop_length * 4
|
| 53 |
+
self.frame_rate = self.sampling_rate / self.inputs_to_logits_ratio
|
| 54 |
+
self.minimum_input_samples = self.n_fft // 2 + 1
|
| 55 |
+
|
| 56 |
+
if min(self.hidden_size, self.intermediate_size, self.num_hidden_layers, self.num_attention_heads) <= 0:
|
| 57 |
+
raise ValueError("Encoder dimensions and layer counts must be positive.")
|
| 58 |
+
if self.hidden_size % self.num_attention_heads or (self.hidden_size // self.num_attention_heads) % 2:
|
| 59 |
+
raise ValueError("hidden_size must divide into an even head dimension.")
|
| 60 |
+
if len(self.subsampling_channels) != 3 or len(self.subsampling_depths) != 3:
|
| 61 |
+
raise ValueError("The subsampler must contain three channel widths and three depths.")
|
| 62 |
+
if self.subsampling_channels[0] != self.num_mel_bins or self.subsampling_channels[-1] != self.hidden_size:
|
| 63 |
+
raise ValueError("Subsampling widths must start at num_mel_bins and end at hidden_size.")
|
| 64 |
+
if min(*self.subsampling_channels, *self.subsampling_depths, self.num_mel_bins, self.sampling_rate, self.hop_length) <= 0:
|
| 65 |
+
raise ValueError("Frontend dimensions and sampling parameters must be positive.")
|
| 66 |
+
if self.n_fft < self.win_length or self.win_length <= 0 or self.n_fft % 2:
|
| 67 |
+
raise ValueError("n_fft must be even and at least win_length > 0.")
|
| 68 |
+
if self.conv_depthwise_kernel_size <= 0 or self.conv_depthwise_kernel_size % 2 != 1:
|
| 69 |
+
raise ValueError("conv_depthwise_kernel_size must be positive and odd.")
|
| 70 |
+
if self.rotary_embedding_base <= 0 or min(self.layer_norm_eps, self.subsampling_layer_norm_eps) <= 0:
|
| 71 |
+
raise ValueError("Rotary base and normalization epsilons must be positive.")
|
| 72 |
+
if self.variant not in {"30s", "fs"}:
|
| 73 |
+
raise ValueError("variant must be '30s' or 'fs'.")
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
MERT2Config.register_for_auto_class()
|
configuration_sheetsage2.py
ADDED
|
@@ -0,0 +1,88 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Configuration for SheetSage2 audio-to-symbolic transcription."""
|
| 2 |
+
|
| 3 |
+
from transformers import PretrainedConfig
|
| 4 |
+
|
| 5 |
+
from .configuration_mert2 import MERT2Config
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
class SheetSage2Config(PretrainedConfig):
|
| 9 |
+
model_type = "sheetsage2"
|
| 10 |
+
|
| 11 |
+
def __init__(
|
| 12 |
+
self,
|
| 13 |
+
vocab_size=31678,
|
| 14 |
+
hidden_size=512,
|
| 15 |
+
decoder_layers=6,
|
| 16 |
+
num_attention_heads=8,
|
| 17 |
+
intermediate_size=2048,
|
| 18 |
+
decoder_dropout=0.1,
|
| 19 |
+
input_audio_length=300.0,
|
| 20 |
+
max_output_seq_len=5120,
|
| 21 |
+
time_hz=100,
|
| 22 |
+
sampling_rate=24000,
|
| 23 |
+
lora_rank=64,
|
| 24 |
+
lora_alpha=128,
|
| 25 |
+
weights_format="adapter",
|
| 26 |
+
backbone_config=None,
|
| 27 |
+
encoder_attn_implementation="sdpa",
|
| 28 |
+
base_model_name_or_path="m-a-p/MERT-v2-FullSong",
|
| 29 |
+
base_model_revision="d8ba1c745e733b3908ce6ad16ebeb17ac7600a42",
|
| 30 |
+
base_model_sha256="e6dd2ab187d6dd62b6521cd7d8f932e237acf0c5757745a7232082e28391350d",
|
| 31 |
+
tokenizer_schema_version="v1",
|
| 32 |
+
tokenizer_fingerprint="5ba3325af0344c7f",
|
| 33 |
+
**kwargs,
|
| 34 |
+
):
|
| 35 |
+
for name, default in (
|
| 36 |
+
("is_encoder_decoder", True), ("tie_word_embeddings", True),
|
| 37 |
+
("pad_token_id", 0), ("bos_token_id", 1), ("eos_token_id", 2),
|
| 38 |
+
("decoder_start_token_id", 1),
|
| 39 |
+
("use_cache", True),
|
| 40 |
+
):
|
| 41 |
+
kwargs.setdefault(name, default)
|
| 42 |
+
super().__init__(**kwargs)
|
| 43 |
+
self.vocab_size = int(vocab_size)
|
| 44 |
+
self.hidden_size = int(hidden_size)
|
| 45 |
+
self.decoder_layers = int(decoder_layers)
|
| 46 |
+
self.num_attention_heads = int(num_attention_heads)
|
| 47 |
+
self.intermediate_size = int(intermediate_size)
|
| 48 |
+
self.decoder_dropout = float(decoder_dropout)
|
| 49 |
+
self.input_audio_length = float(input_audio_length)
|
| 50 |
+
self.max_output_seq_len = int(max_output_seq_len)
|
| 51 |
+
self.time_hz = int(time_hz)
|
| 52 |
+
self.sampling_rate = int(sampling_rate)
|
| 53 |
+
self.lora_rank = int(lora_rank)
|
| 54 |
+
self.lora_alpha = float(lora_alpha)
|
| 55 |
+
self.weights_format = str(weights_format)
|
| 56 |
+
self.backbone_config = dict(backbone_config or MERT2Config(variant="fs").to_dict())
|
| 57 |
+
self.backbone_config.pop("_name_or_path", None)
|
| 58 |
+
self.backbone_config.pop("auto_map", None)
|
| 59 |
+
self.encoder_attn_implementation = str(encoder_attn_implementation)
|
| 60 |
+
self.base_model_name_or_path = str(base_model_name_or_path)
|
| 61 |
+
self.base_model_revision = str(base_model_revision)
|
| 62 |
+
self.base_model_sha256 = str(base_model_sha256)
|
| 63 |
+
self.tokenizer_schema_version = str(tokenizer_schema_version)
|
| 64 |
+
self.tokenizer_fingerprint = str(tokenizer_fingerprint)
|
| 65 |
+
self.architectures = ["SheetSage2Model"]
|
| 66 |
+
self.auto_map = {
|
| 67 |
+
"AutoConfig": "configuration_sheetsage2.SheetSage2Config",
|
| 68 |
+
"AutoModel": "modeling_sheetsage2.SheetSage2Model",
|
| 69 |
+
"AutoModelForSeq2SeqLM": "modeling_sheetsage2.SheetSage2Model",
|
| 70 |
+
"AutoProcessor": "processing_sheetsage2.SheetSage2Processor",
|
| 71 |
+
}
|
| 72 |
+
if self.weights_format not in {"adapter", "merged"}:
|
| 73 |
+
raise ValueError("weights_format must be 'adapter' or 'merged'.")
|
| 74 |
+
if self.encoder_attn_implementation not in {"sdpa", "flash_attention_2"}:
|
| 75 |
+
raise ValueError("Select encoder attention 'sdpa' or 'flash_attention_2'.")
|
| 76 |
+
if min(self.vocab_size, self.hidden_size, self.decoder_layers, self.num_attention_heads,
|
| 77 |
+
self.intermediate_size, self.input_audio_length, self.max_output_seq_len,
|
| 78 |
+
self.time_hz, self.sampling_rate, self.lora_rank, self.lora_alpha) <= 0:
|
| 79 |
+
raise ValueError("Model dimensions and timing parameters must be positive.")
|
| 80 |
+
if self.hidden_size % self.num_attention_heads:
|
| 81 |
+
raise ValueError("hidden_size must be divisible by num_attention_heads.")
|
| 82 |
+
if not 0 <= self.decoder_dropout < 1:
|
| 83 |
+
raise ValueError("decoder_dropout must lie in [0, 1).")
|
| 84 |
+
if self.sampling_rate != self.backbone_config["sampling_rate"]:
|
| 85 |
+
raise ValueError("Processor and encoder sampling rates must match.")
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
SheetSage2Config.register_for_auto_class()
|
durations_sheetsage2.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import numpy as np
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
DURATION_TEMPLATES = np.array(
|
| 5 |
+
[
|
| 6 |
+
1, 2, 3, 4, 6, 8, 12, 16, 24, 32, 48, 64,
|
| 7 |
+
96, 128, 192, 256, 384, 512, 768, 1024,
|
| 8 |
+
1536, 2048, 3072, 4096,
|
| 9 |
+
],
|
| 10 |
+
dtype=np.int32,
|
| 11 |
+
)
|
| 12 |
+
duration_boundaries = (DURATION_TEMPLATES[:-1] + DURATION_TEMPLATES[1:]) / 2
|
| 13 |
+
|
| 14 |
+
|
exports_sheetsage2.py
ADDED
|
@@ -0,0 +1,200 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Lossless decoded events/LAB/MIDI, then independently validated ABC."""
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
import json
|
| 4 |
+
from fractions import Fraction
|
| 5 |
+
|
| 6 |
+
import numpy as np
|
| 7 |
+
import pretty_midi
|
| 8 |
+
|
| 9 |
+
from .notation_sheetsage2 import generate_abc_from_data
|
| 10 |
+
from .io_sheetsage2 import atomic_write_text
|
| 11 |
+
from .generation_sheetsage2 import field_text
|
| 12 |
+
from .midi_sheetsage2 import build_playback, midi_bytes
|
| 13 |
+
|
| 14 |
+
|
| 15 |
+
def rows_text(rows):
|
| 16 |
+
return "".join("\t".join(str(x) for x in row) + "\n" for row in rows)
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def write_rows(path, rows):
|
| 20 |
+
atomic_write_text(path, rows_text(rows))
|
| 21 |
+
|
| 22 |
+
|
| 23 |
+
def interval_rows(events, field, duration):
|
| 24 |
+
rows = [[e["time"], 0, e["values"][field]] for e in events if field in e["values"]]
|
| 25 |
+
for i, row in enumerate(rows):
|
| 26 |
+
row[1] = rows[i + 1][0] if i + 1 < len(rows) else duration
|
| 27 |
+
return [r for r in rows if r[1] > r[0]]
|
| 28 |
+
|
| 29 |
+
|
| 30 |
+
def rhythm_rows(events):
|
| 31 |
+
rows, meter = [], None
|
| 32 |
+
for event in events:
|
| 33 |
+
rhythm = event["values"].get("rhythm", {})
|
| 34 |
+
meter = rhythm.get("meter", meter)
|
| 35 |
+
eighth = rhythm.get("eighth_position")
|
| 36 |
+
if eighth is not None and meter is not None:
|
| 37 |
+
# The model stores eighth-note position, not denominator-beat index.
|
| 38 |
+
position = Fraction(int(eighth) * int(meter[1]), 8)
|
| 39 |
+
if position.denominator != 1:
|
| 40 |
+
raise ValueError(f"Eighth position {eighth} is off the {meter[0]}/{meter[1]} beat grid")
|
| 41 |
+
if not 0 <= position < meter[0]:
|
| 42 |
+
raise ValueError(f"Eighth position {eighth} is outside meter {meter}")
|
| 43 |
+
rows.append([float(event["time"]), int(position) + 1, int(meter[0]), int(meter[1])])
|
| 44 |
+
return rows
|
| 45 |
+
|
| 46 |
+
|
| 47 |
+
def _midi(notes, paper=False):
|
| 48 |
+
midi = pretty_midi.PrettyMIDI(resolution=220 if paper else 960)
|
| 49 |
+
names = ((0, "vocal_melody"), (1, "instrumental_melody")) if paper else ((0, "Vocal"), (1, "Ins"))
|
| 50 |
+
for track, name in names:
|
| 51 |
+
instrument = pretty_midi.Instrument(program=0, name=name)
|
| 52 |
+
for start, end, pitch, source in notes:
|
| 53 |
+
if track == source and end > start:
|
| 54 |
+
instrument.notes.append(pretty_midi.Note(100, pitch, start, end))
|
| 55 |
+
midi.instruments.append(instrument)
|
| 56 |
+
return midi
|
| 57 |
+
|
| 58 |
+
|
| 59 |
+
def notation_notes(notes):
|
| 60 |
+
"""A monophonic notation view, retaining the exact raw prediction separately."""
|
| 61 |
+
result, diagnostics = [], []
|
| 62 |
+
for track in (0, 1):
|
| 63 |
+
ordered = sorted((list(n) for n in notes if n[3] == track), key=lambda n: (n[0], n[2], n[1]))
|
| 64 |
+
for i, note in enumerate(ordered):
|
| 65 |
+
if i + 1 < len(ordered) and note[1] > ordered[i + 1][0] + 1e-6:
|
| 66 |
+
note[1] = ordered[i + 1][0]
|
| 67 |
+
diagnostics.append(f"notation only: clipped track {track} note at {note[0]:.3f} to next onset")
|
| 68 |
+
if note[1] > note[0] + 1e-6:
|
| 69 |
+
result.append(note)
|
| 70 |
+
return sorted(result), diagnostics
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
def export_result(decoded, tokenizer, output_dir=None, duration=None, paper=False, *, melody_only=False):
|
| 74 |
+
"""Build ABC, MIDI and LAB in memory, optionally writing the same bytes.
|
| 75 |
+
|
| 76 |
+
Statistics retain their file-interface names. ``payload`` contains the
|
| 77 |
+
in-memory ABC text, MIDI bytes, decoded event list, LAB texts and playback.
|
| 78 |
+
melody_only removes chords from notation and playback without changing
|
| 79 |
+
decoded events or raw LAB annotations.
|
| 80 |
+
"""
|
| 81 |
+
if not isinstance(melody_only, bool):
|
| 82 |
+
raise ValueError("melody_only must be True or False")
|
| 83 |
+
if duration is None:
|
| 84 |
+
raise TypeError("duration is required")
|
| 85 |
+
texts, midis, notation_midis, labs = {}, {}, {}, {}
|
| 86 |
+
|
| 87 |
+
def keep_rows(name, rows):
|
| 88 |
+
text = rows_text(rows)
|
| 89 |
+
texts[name] = text
|
| 90 |
+
if name.endswith(".lab"):
|
| 91 |
+
labs[name[:-4]] = text
|
| 92 |
+
|
| 93 |
+
events = decoded["events"]
|
| 94 |
+
notes = []
|
| 95 |
+
for event in events:
|
| 96 |
+
start = float(event["time"])
|
| 97 |
+
for note in event["values"].get("melody", ()):
|
| 98 |
+
end = max(start + 0.04, float(note["end_time"])) if paper else min(duration, float(note["end_time"]))
|
| 99 |
+
if end > start:
|
| 100 |
+
notes.append([start, end, int(note["pitch"]), int(note["track"])])
|
| 101 |
+
if not paper:
|
| 102 |
+
notes.sort()
|
| 103 |
+
texts["events.json"] = json.dumps(decoded, indent=2)
|
| 104 |
+
keep_rows("events.tsv", [["time", "global_subbeat", "fields"]] +
|
| 105 |
+
[[e["time"], e["global_subbeat"], field_text(e, tokenizer)] for e in events])
|
| 106 |
+
midis["melody"] = midi_bytes(_midi(notes, paper=paper))
|
| 107 |
+
lab_notes = [[f"{a:.6f}", f"{b:.6f}", p, t] for a, b, p, t in notes] if paper else notes
|
| 108 |
+
keep_rows("melody_full.lab", lab_notes)
|
| 109 |
+
for track, name in ((0, "vocal"), (1, "instrumental")):
|
| 110 |
+
selected = [n for n in notes if n[3] == track]
|
| 111 |
+
keep_rows(f"melody_{name}.lab", [n[:3] for n in lab_notes if n[3] == track])
|
| 112 |
+
midis[f"melody_{name}"] = midi_bytes(_midi(selected, paper=paper))
|
| 113 |
+
intervals = {}
|
| 114 |
+
for field in ("chord", "key", "structure"):
|
| 115 |
+
rows = interval_rows(events, field, duration)
|
| 116 |
+
if paper:
|
| 117 |
+
rows = [[float(e["time"]), None, e["values"][field]] for e in events if field in e["values"]]
|
| 118 |
+
for i, row in enumerate(rows):
|
| 119 |
+
row[1] = max(row[0], rows[i + 1][0]) if i + 1 < len(rows) else float(duration)
|
| 120 |
+
intervals[field] = rows
|
| 121 |
+
keep_rows(f"{field}.lab", rows)
|
| 122 |
+
if paper:
|
| 123 |
+
rows = []
|
| 124 |
+
for event in events:
|
| 125 |
+
rhythm, stamp = event["values"].get("rhythm"), event["values"].get("timestamp")
|
| 126 |
+
if rhythm is None and stamp is None:
|
| 127 |
+
continue
|
| 128 |
+
meter = rhythm.get("meter") if isinstance(rhythm, dict) else None
|
| 129 |
+
eighth = rhythm.get("eighth_position") if isinstance(rhythm, dict) else None
|
| 130 |
+
rows.append([f"{event['time']:.6f}", "" if stamp is None else f"{float(stamp):.6f}",
|
| 131 |
+
"" if eighth is None else int(eighth), "" if meter is None else f"{meter[0]}/{meter[1]}"])
|
| 132 |
+
keep_rows("beat_meter.lab", rows)
|
| 133 |
+
raw_rhythm = [[e["time"], json.dumps(e["values"].get("rhythm", {}))]
|
| 134 |
+
for e in events if "rhythm" in e["values"] or "timestamp" in e["values"]]
|
| 135 |
+
keep_rows("rhythm_events.lab", raw_rhythm)
|
| 136 |
+
diagnostics, abc_error, score, text = [], None, None, None
|
| 137 |
+
try:
|
| 138 |
+
beats = rhythm_rows(events)
|
| 139 |
+
keep_rows("beat.lab", beats)
|
| 140 |
+
keep_rows("downbeat.lab", [[r[0]] for r in beats if r[1] == 1])
|
| 141 |
+
if len(beats) < 2:
|
| 142 |
+
raise ValueError("At least two decoded beats are required for ABC")
|
| 143 |
+
# canonical notation's last beat is a boundary, so append it explicitly. Continue the
|
| 144 |
+
# final tempo only as far as the audio/last note and let it pad the bar.
|
| 145 |
+
abc_beats = [list(b) for b in beats]
|
| 146 |
+
period = float(np.median(np.diff([b[0] for b in beats[-9:]])))
|
| 147 |
+
if period <= 0:
|
| 148 |
+
raise ValueError("Decoded beats must increase in time")
|
| 149 |
+
end = max(duration, max((n[1] for n in notes), default=0))
|
| 150 |
+
while abc_beats[-1][0] < end - 1e-6:
|
| 151 |
+
prev = abc_beats[-1]
|
| 152 |
+
abc_beats.append([prev[0] + period, prev[1] % prev[2] + 1, prev[2], prev[3]])
|
| 153 |
+
keep_rows("notation/song_beats.txt", abc_beats)
|
| 154 |
+
notation_intervals = {}
|
| 155 |
+
for field, suffix in (("chord", "chords"), ("key", "keys"), ("structure", "structures")):
|
| 156 |
+
rows = interval_rows(events, field, duration)
|
| 157 |
+
# Clip to the actual beat domain: events beyond it have no ABC bin.
|
| 158 |
+
rows = [[max(abc_beats[0][0], a), min(abc_beats[-1][0], b), v]
|
| 159 |
+
for a, b, v in rows if b > abc_beats[0][0] and a < abc_beats[-1][0]]
|
| 160 |
+
notation_intervals[field] = rows
|
| 161 |
+
keep_rows(f"notation/song_{suffix}.txt", rows)
|
| 162 |
+
if not interval_rows(events, "key", duration):
|
| 163 |
+
raise ValueError("No key was decoded; cannot construct a keyed ABC score")
|
| 164 |
+
clean, adjustments = notation_notes(notes)
|
| 165 |
+
diagnostics.extend(adjustments)
|
| 166 |
+
notation_midis["notation/song_melody.mid"] = midi_bytes(_midi(clean))
|
| 167 |
+
text, score = generate_abc_from_data(
|
| 168 |
+
notation_midis["notation/song_melody.mid"], abc_beats,
|
| 169 |
+
notation_intervals["chord"], notation_intervals["key"], notation_intervals["structure"],
|
| 170 |
+
melody_only=melody_only,
|
| 171 |
+
)
|
| 172 |
+
diagnostics.extend(score.diagnostics)
|
| 173 |
+
texts["score.abc"] = text
|
| 174 |
+
abc_measures = len(score.measures)
|
| 175 |
+
except (ValueError, FileNotFoundError) as exc:
|
| 176 |
+
abc_error = str(exc)
|
| 177 |
+
abc_measures = 0
|
| 178 |
+
playback, playback_midis = build_playback(
|
| 179 |
+
midis["melody"], [] if melody_only else intervals["chord"], score, duration)
|
| 180 |
+
midis.update(playback_midis)
|
| 181 |
+
texts["playback.json"] = json.dumps(playback, indent=2)
|
| 182 |
+
diagnostics.extend(playback["warnings"])
|
| 183 |
+
if output_dir is not None:
|
| 184 |
+
output_dir = Path(output_dir)
|
| 185 |
+
output_dir.mkdir(parents=True, exist_ok=True)
|
| 186 |
+
for name, content in texts.items():
|
| 187 |
+
atomic_write_text(output_dir / name, content)
|
| 188 |
+
for name, content in {**{f"{name}.mid": content for name, content in midis.items()}, **notation_midis}.items():
|
| 189 |
+
path = output_dir / name
|
| 190 |
+
path.parent.mkdir(parents=True, exist_ok=True)
|
| 191 |
+
path.write_bytes(content)
|
| 192 |
+
if abc_error is not None:
|
| 193 |
+
(output_dir / "score.abc").unlink(missing_ok=True)
|
| 194 |
+
else:
|
| 195 |
+
(output_dir / "score.browser.abc").unlink(missing_ok=True)
|
| 196 |
+
payload = dict(abc=text, midi=midis["transcription"], midis=midis,
|
| 197 |
+
events=events, labs=labs, playback=playback)
|
| 198 |
+
return dict(melody_notes=len(notes), vocal_notes=sum(n[3] == 0 for n in notes),
|
| 199 |
+
instrumental_notes=sum(n[3] == 1 for n in notes), events=len(events),
|
| 200 |
+
abc_measures=abc_measures, abc_error=abc_error, diagnostics=diagnostics, payload=payload)
|
generation_sheetsage2.py
ADDED
|
@@ -0,0 +1,634 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Prompt grammar, cached generation and overlap context from the evaluated model."""
|
| 2 |
+
import copy
|
| 3 |
+
from contextlib import nullcontext
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import numpy as np
|
| 6 |
+
import torch
|
| 7 |
+
|
| 8 |
+
FULL_TASK_PROMPTS = ("timestamp", "downbeat_meter", "structure", "key", "chord_full", "melody_full")
|
| 9 |
+
FIELD_TO_INDEX = {"timestamp": 0, "rhythm": 1, "structure": 2, "key": 3, "chord": 4, "melody": 5}
|
| 10 |
+
|
| 11 |
+
|
| 12 |
+
class PromptGrammarState:
|
| 13 |
+
def __init__(self, tokenizer):
|
| 14 |
+
self.tokenizer = tokenizer
|
| 15 |
+
self.generated_events = 0
|
| 16 |
+
self.in_shift = True
|
| 17 |
+
self.shift_run = 0
|
| 18 |
+
self.payload_count = 0
|
| 19 |
+
self.last_field_index = -1
|
| 20 |
+
self.incomplete = None
|
| 21 |
+
|
| 22 |
+
def _allow_field_starts(self, allowed):
|
| 23 |
+
tokenizer = self.tokenizer
|
| 24 |
+
if self.last_field_index < FIELD_TO_INDEX["timestamp"]:
|
| 25 |
+
allowed[tokenizer.time_token_start : tokenizer.time_token_end] = True
|
| 26 |
+
if self.last_field_index < FIELD_TO_INDEX["rhythm"]:
|
| 27 |
+
allowed[tokenizer.meter_token_start : tokenizer.meter_token_end] = True
|
| 28 |
+
allowed[
|
| 29 |
+
tokenizer.eighth_position_token_start : tokenizer.eighth_position_token_end
|
| 30 |
+
] = True
|
| 31 |
+
if self.last_field_index < FIELD_TO_INDEX["structure"]:
|
| 32 |
+
allowed[
|
| 33 |
+
tokenizer.structure_token_start : tokenizer.structure_token_end
|
| 34 |
+
] = True
|
| 35 |
+
if self.last_field_index < FIELD_TO_INDEX["key"]:
|
| 36 |
+
allowed[tokenizer.key_token_start : tokenizer.key_token_end] = True
|
| 37 |
+
if self.last_field_index < FIELD_TO_INDEX["chord"]:
|
| 38 |
+
allowed[
|
| 39 |
+
tokenizer.full_chord_token_start : tokenizer.full_chord_token_end
|
| 40 |
+
] = True
|
| 41 |
+
if self.last_field_index <= FIELD_TO_INDEX["melody"]:
|
| 42 |
+
allowed[tokenizer.pitch_token_start : tokenizer.pitch_token_end] = True
|
| 43 |
+
|
| 44 |
+
def allowed(self, device):
|
| 45 |
+
tokenizer = self.tokenizer
|
| 46 |
+
allowed = torch.zeros(tokenizer.n_tokens, dtype=torch.bool, device=device)
|
| 47 |
+
can_end = self.payload_count > 0
|
| 48 |
+
|
| 49 |
+
if can_end:
|
| 50 |
+
allowed[tokenizer.eos_token] = True
|
| 51 |
+
if self.payload_count > 0 or self.in_shift:
|
| 52 |
+
if self.shift_run < 4:
|
| 53 |
+
allowed[
|
| 54 |
+
tokenizer.subbeat_shift_token_start : tokenizer.subbeat_shift_token_end
|
| 55 |
+
] = True
|
| 56 |
+
|
| 57 |
+
if self.incomplete == "rhythm_after_meter":
|
| 58 |
+
allowed[
|
| 59 |
+
tokenizer.eighth_position_token_start : tokenizer.eighth_position_token_end
|
| 60 |
+
] = True
|
| 61 |
+
return allowed
|
| 62 |
+
|
| 63 |
+
if self.incomplete == "melody_after_pitch":
|
| 64 |
+
allowed[tokenizer.duration_token_start : tokenizer.duration_token_end] = True
|
| 65 |
+
allowed[tokenizer.pitch_token_start : tokenizer.pitch_token_end] = True
|
| 66 |
+
return allowed
|
| 67 |
+
|
| 68 |
+
self._allow_field_starts(allowed)
|
| 69 |
+
return allowed
|
| 70 |
+
|
| 71 |
+
def update(self, token):
|
| 72 |
+
tokenizer = self.tokenizer
|
| 73 |
+
token = int(token)
|
| 74 |
+
token_type = tokenizer.token_type(token)
|
| 75 |
+
if token == tokenizer.eos_token:
|
| 76 |
+
return True
|
| 77 |
+
if token_type == "subbeat_shift":
|
| 78 |
+
if not self.in_shift and self.payload_count > 0:
|
| 79 |
+
self.generated_events += 1
|
| 80 |
+
self.payload_count = 0
|
| 81 |
+
self.last_field_index = -1
|
| 82 |
+
self.incomplete = None
|
| 83 |
+
self.in_shift = True
|
| 84 |
+
self.shift_run += 1
|
| 85 |
+
return False
|
| 86 |
+
|
| 87 |
+
self.in_shift = False
|
| 88 |
+
self.shift_run = 0
|
| 89 |
+
self.payload_count += 1
|
| 90 |
+
if token_type == "time":
|
| 91 |
+
self.last_field_index = FIELD_TO_INDEX["timestamp"]
|
| 92 |
+
self.incomplete = None
|
| 93 |
+
elif token_type == "meter":
|
| 94 |
+
self.last_field_index = FIELD_TO_INDEX["rhythm"]
|
| 95 |
+
self.incomplete = "rhythm_after_meter"
|
| 96 |
+
elif token_type == "eighth_position":
|
| 97 |
+
self.last_field_index = FIELD_TO_INDEX["rhythm"]
|
| 98 |
+
self.incomplete = None
|
| 99 |
+
elif token_type == "structure":
|
| 100 |
+
self.last_field_index = FIELD_TO_INDEX["structure"]
|
| 101 |
+
self.incomplete = None
|
| 102 |
+
elif token_type == "key":
|
| 103 |
+
self.last_field_index = FIELD_TO_INDEX["key"]
|
| 104 |
+
self.incomplete = None
|
| 105 |
+
elif token_type == "chord_full":
|
| 106 |
+
self.last_field_index = FIELD_TO_INDEX["chord"]
|
| 107 |
+
self.incomplete = None
|
| 108 |
+
elif token_type == "pitch":
|
| 109 |
+
self.last_field_index = FIELD_TO_INDEX["melody"]
|
| 110 |
+
self.incomplete = "melody_after_pitch"
|
| 111 |
+
elif token_type == "duration":
|
| 112 |
+
self.last_field_index = FIELD_TO_INDEX["melody"]
|
| 113 |
+
self.incomplete = None
|
| 114 |
+
else:
|
| 115 |
+
raise RuntimeError(f"Unexpected prompt token type {token_type!r}")
|
| 116 |
+
return False
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
def inference_autocast(device, dtype=torch.bfloat16):
|
| 120 |
+
if device.type != "cuda" or dtype is None:
|
| 121 |
+
return nullcontext()
|
| 122 |
+
return torch.autocast(device_type="cuda", dtype=dtype)
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
def select_cache_batch(cache, indices, batch_size):
|
| 126 |
+
if cache is None:
|
| 127 |
+
return None
|
| 128 |
+
batch_select = getattr(cache, "batch_select_indices", None)
|
| 129 |
+
if callable(batch_select):
|
| 130 |
+
batch_select(indices)
|
| 131 |
+
return cache
|
| 132 |
+
if torch.is_tensor(cache):
|
| 133 |
+
if cache.ndim > 0 and cache.shape[0] == batch_size:
|
| 134 |
+
return cache.index_select(0, indices)
|
| 135 |
+
return cache
|
| 136 |
+
if isinstance(cache, tuple):
|
| 137 |
+
return tuple(select_cache_batch(item, indices, batch_size) for item in cache)
|
| 138 |
+
if isinstance(cache, list):
|
| 139 |
+
return [select_cache_batch(item, indices, batch_size) for item in cache]
|
| 140 |
+
if isinstance(cache, dict):
|
| 141 |
+
return {
|
| 142 |
+
key: select_cache_batch(value, indices, batch_size)
|
| 143 |
+
for key, value in cache.items()
|
| 144 |
+
}
|
| 145 |
+
return cache
|
| 146 |
+
|
| 147 |
+
|
| 148 |
+
@torch.inference_mode()
|
| 149 |
+
def constrained_prompt_generate_batch(
|
| 150 |
+
model,
|
| 151 |
+
audio,
|
| 152 |
+
prompts,
|
| 153 |
+
max_sequence_length,
|
| 154 |
+
prefix_tokens=None,
|
| 155 |
+
autocast_dtype=torch.bfloat16,
|
| 156 |
+
stop_time_seconds=None,
|
| 157 |
+
progress_callback=None,
|
| 158 |
+
memory=None,
|
| 159 |
+
step_callback=None,
|
| 160 |
+
):
|
| 161 |
+
tokenizer = model.tokenizer
|
| 162 |
+
prefix = list(prefix_tokens) if prefix_tokens is not None else tokenizer.prompt_prefix(prompts)
|
| 163 |
+
if not prefix or prefix[0] != tokenizer.sos_token:
|
| 164 |
+
raise ValueError("generation prefix must begin with <|sos|>")
|
| 165 |
+
if prefix[-1] == tokenizer.eos_token:
|
| 166 |
+
prefix = prefix[:-1]
|
| 167 |
+
if audio.ndim != 2 or audio.shape[0] < 1:
|
| 168 |
+
raise ValueError(f"audio must have shape [batch, samples], got {tuple(audio.shape)}")
|
| 169 |
+
states = [PromptGrammarState(tokenizer) for _ in range(audio.shape[0])]
|
| 170 |
+
stop_times = [
|
| 171 |
+
None if stop_time_seconds is None else float(stop_time_seconds)
|
| 172 |
+
for _ in states
|
| 173 |
+
]
|
| 174 |
+
out_index = prefix.index(tokenizer.out_token)
|
| 175 |
+
for state in states:
|
| 176 |
+
for token in prefix[out_index + 1 :]:
|
| 177 |
+
state.update(token)
|
| 178 |
+
output_tokens = [list(prefix) for _ in states]
|
| 179 |
+
active_sample_ids = list(range(audio.shape[0]))
|
| 180 |
+
decoder_input = torch.tensor(
|
| 181 |
+
[prefix] * audio.shape[0],
|
| 182 |
+
dtype=torch.long,
|
| 183 |
+
device=audio.device,
|
| 184 |
+
)
|
| 185 |
+
past_key_values = None
|
| 186 |
+
if memory is None:
|
| 187 |
+
with inference_autocast(audio.device, autocast_dtype):
|
| 188 |
+
memory = model.encode(audio)
|
| 189 |
+
current_length = len(prefix)
|
| 190 |
+
while active_sample_ids and current_length < int(max_sequence_length):
|
| 191 |
+
with inference_autocast(audio.device, autocast_dtype):
|
| 192 |
+
logits, past_key_values = model.decode(
|
| 193 |
+
memory,
|
| 194 |
+
decoder_input,
|
| 195 |
+
use_cache=True,
|
| 196 |
+
past_key_values=past_key_values,
|
| 197 |
+
)
|
| 198 |
+
next_logits = logits[:, -1].float()
|
| 199 |
+
allowed = torch.stack(
|
| 200 |
+
[state.allowed(next_logits.device) for state in states],
|
| 201 |
+
dim=0,
|
| 202 |
+
)
|
| 203 |
+
masked_logits = next_logits.masked_fill(~allowed, float("-inf"))
|
| 204 |
+
next_tokens = masked_logits.argmax(dim=-1)
|
| 205 |
+
if step_callback is not None:
|
| 206 |
+
step_callback(current_length, active_sample_ids, next_logits, masked_logits)
|
| 207 |
+
next_token_values = next_tokens.tolist()
|
| 208 |
+
keep_positions = []
|
| 209 |
+
for position, (sample_id, state, stop_time, token) in enumerate(
|
| 210 |
+
zip(active_sample_ids, states, stop_times, next_token_values)
|
| 211 |
+
):
|
| 212 |
+
output_tokens[sample_id].append(int(token))
|
| 213 |
+
finished = state.update(token)
|
| 214 |
+
if (
|
| 215 |
+
not finished
|
| 216 |
+
and stop_time is not None
|
| 217 |
+
and tokenizer.time_token_start <= int(token) < tokenizer.time_token_end
|
| 218 |
+
and tokenizer.token_to_time_id(token) / tokenizer.time_hz >= stop_time
|
| 219 |
+
):
|
| 220 |
+
output_tokens[sample_id].append(tokenizer.eos_token)
|
| 221 |
+
finished = True
|
| 222 |
+
if not finished:
|
| 223 |
+
keep_positions.append(position)
|
| 224 |
+
current_length += 1
|
| 225 |
+
if progress_callback is not None and current_length % 64 == 0:
|
| 226 |
+
progress_callback(current_length)
|
| 227 |
+
if not keep_positions:
|
| 228 |
+
break
|
| 229 |
+
|
| 230 |
+
old_batch_size = len(active_sample_ids)
|
| 231 |
+
decoder_input = next_tokens[:, None]
|
| 232 |
+
if len(keep_positions) == old_batch_size:
|
| 233 |
+
continue
|
| 234 |
+
|
| 235 |
+
keep = torch.tensor(keep_positions, dtype=torch.long, device=audio.device)
|
| 236 |
+
active_sample_ids = [active_sample_ids[position] for position in keep_positions]
|
| 237 |
+
states = [states[position] for position in keep_positions]
|
| 238 |
+
stop_times = [stop_times[position] for position in keep_positions]
|
| 239 |
+
memory = memory.index_select(0, keep)
|
| 240 |
+
past_key_values = select_cache_batch(past_key_values, keep, old_batch_size)
|
| 241 |
+
decoder_input = decoder_input.index_select(0, keep)
|
| 242 |
+
|
| 243 |
+
for tokens in output_tokens:
|
| 244 |
+
if tokens[-1] != tokenizer.eos_token:
|
| 245 |
+
tokens.append(tokenizer.eos_token)
|
| 246 |
+
return [torch.tensor(tokens, dtype=torch.long) for tokens in output_tokens]
|
| 247 |
+
|
| 248 |
+
|
| 249 |
+
@torch.inference_mode()
|
| 250 |
+
def constrained_prompt_generate(
|
| 251 |
+
model,
|
| 252 |
+
audio,
|
| 253 |
+
prompts,
|
| 254 |
+
max_sequence_length,
|
| 255 |
+
prefix_tokens=None,
|
| 256 |
+
autocast_dtype=torch.bfloat16,
|
| 257 |
+
stop_time_seconds=None,
|
| 258 |
+
progress_callback=None,
|
| 259 |
+
memory=None,
|
| 260 |
+
step_callback=None,
|
| 261 |
+
):
|
| 262 |
+
if audio.shape[0] != 1:
|
| 263 |
+
raise ValueError(
|
| 264 |
+
"constrained_prompt_generate expects batch size 1; "
|
| 265 |
+
"use constrained_prompt_generate_batch for batched inference"
|
| 266 |
+
)
|
| 267 |
+
return constrained_prompt_generate_batch(
|
| 268 |
+
model,
|
| 269 |
+
audio,
|
| 270 |
+
prompts,
|
| 271 |
+
max_sequence_length,
|
| 272 |
+
prefix_tokens=prefix_tokens,
|
| 273 |
+
autocast_dtype=autocast_dtype,
|
| 274 |
+
stop_time_seconds=stop_time_seconds,
|
| 275 |
+
progress_callback=progress_callback,
|
| 276 |
+
memory=memory,
|
| 277 |
+
step_callback=step_callback,
|
| 278 |
+
)[0]
|
| 279 |
+
|
| 280 |
+
|
| 281 |
+
def decode_generated_tokens(tokenizer, tokens, song_id, window_index=None):
|
| 282 |
+
try:
|
| 283 |
+
return tokenizer.decode_sequence(tokens, strict=True), None
|
| 284 |
+
except ValueError as exc:
|
| 285 |
+
recoverable_errors = (
|
| 286 |
+
"empty event at subbeat",
|
| 287 |
+
"belongs to inactive output field",
|
| 288 |
+
)
|
| 289 |
+
if not any(message in str(exc) for message in recoverable_errors):
|
| 290 |
+
raise
|
| 291 |
+
decoded = tokenizer.decode_sequence(tokens, strict=False)
|
| 292 |
+
warning = {
|
| 293 |
+
"song_id": song_id,
|
| 294 |
+
"warning": "strict decode failed; recovered in non-strict decode",
|
| 295 |
+
"error": str(exc),
|
| 296 |
+
}
|
| 297 |
+
if window_index is not None:
|
| 298 |
+
warning["window_index"] = int(window_index)
|
| 299 |
+
return decoded, warning
|
| 300 |
+
|
| 301 |
+
|
| 302 |
+
def event_time_map(decoded, target_seconds):
|
| 303 |
+
anchors = []
|
| 304 |
+
for event in decoded["events"]:
|
| 305 |
+
value = event["values"].get("timestamp")
|
| 306 |
+
if value is not None:
|
| 307 |
+
anchors.append((int(event["subbeat"]), float(value)))
|
| 308 |
+
if not anchors:
|
| 309 |
+
return lambda step: min(float(target_seconds), max(0.0, float(step) * 0.125))
|
| 310 |
+
anchors = sorted(dict(anchors).items())
|
| 311 |
+
steps = np.asarray([item[0] for item in anchors], dtype=np.float64)
|
| 312 |
+
times = np.asarray([item[1] for item in anchors], dtype=np.float64)
|
| 313 |
+
if len(anchors) >= 2:
|
| 314 |
+
step_seconds = float(np.median(np.diff(times) / np.maximum(np.diff(steps), 1)))
|
| 315 |
+
if not np.isfinite(step_seconds) or step_seconds <= 0:
|
| 316 |
+
step_seconds = 0.125
|
| 317 |
+
else:
|
| 318 |
+
step_seconds = 0.125
|
| 319 |
+
|
| 320 |
+
def lookup(step):
|
| 321 |
+
step = float(step)
|
| 322 |
+
if step <= steps[0]:
|
| 323 |
+
return float(np.clip(times[0] + (step - steps[0]) * step_seconds, 0, target_seconds))
|
| 324 |
+
if step >= steps[-1]:
|
| 325 |
+
return float(np.clip(times[-1] + (step - steps[-1]) * step_seconds, 0, target_seconds))
|
| 326 |
+
return float(np.interp(step, steps, times))
|
| 327 |
+
|
| 328 |
+
return lookup
|
| 329 |
+
|
| 330 |
+
|
| 331 |
+
def field_text(event, tokenizer):
|
| 332 |
+
parts = []
|
| 333 |
+
for field in tokenizer.event_field_order:
|
| 334 |
+
value = event["values"].get(field)
|
| 335 |
+
if value is None:
|
| 336 |
+
continue
|
| 337 |
+
if field == "melody":
|
| 338 |
+
note_parts = []
|
| 339 |
+
for note in value:
|
| 340 |
+
duration = note["duration_bin"]
|
| 341 |
+
note_parts.append(
|
| 342 |
+
f"pitch={note['pitch']}:track={note['track']}:dur_bin={duration}:dur_steps={note['duration_steps']}"
|
| 343 |
+
)
|
| 344 |
+
parts.append("melody=[" + ",".join(note_parts) + "]")
|
| 345 |
+
elif isinstance(value, dict):
|
| 346 |
+
parts.append(field + "=" + ",".join(f"{k}:{v}" for k, v in value.items()))
|
| 347 |
+
else:
|
| 348 |
+
parts.append(f"{field}={value}")
|
| 349 |
+
return "; ".join(parts)
|
| 350 |
+
|
| 351 |
+
|
| 352 |
+
def event_start_time(event, time_lookup):
|
| 353 |
+
if "time" in event:
|
| 354 |
+
return float(event["time"])
|
| 355 |
+
return float(time_lookup(event["subbeat"]))
|
| 356 |
+
|
| 357 |
+
|
| 358 |
+
def write_events_tsv(decoded, tokenizer, time_lookup, path):
|
| 359 |
+
with Path(path).open("w", encoding="utf-8") as f:
|
| 360 |
+
f.write("event_index\tsubbeat\ttime\tfields\ttokens\n")
|
| 361 |
+
for index, event in enumerate(decoded["events"]):
|
| 362 |
+
token_text = " ".join(
|
| 363 |
+
tokenizer.describe(token)
|
| 364 |
+
for field in tokenizer.event_field_order
|
| 365 |
+
for token in event["tokens_by_field"].get(field, ())
|
| 366 |
+
)
|
| 367 |
+
f.write(
|
| 368 |
+
f"{index}\t{event['subbeat']}\t{event_start_time(event, time_lookup):.6f}\t"
|
| 369 |
+
f"{field_text(event, tokenizer)}\t{token_text}\n"
|
| 370 |
+
)
|
| 371 |
+
|
| 372 |
+
|
| 373 |
+
def clone_decoded_event(event):
|
| 374 |
+
return {
|
| 375 |
+
"subbeat": int(event["subbeat"]),
|
| 376 |
+
"tokens_by_field": {
|
| 377 |
+
field: [int(token) for token in tokens]
|
| 378 |
+
for field, tokens in event["tokens_by_field"].items()
|
| 379 |
+
},
|
| 380 |
+
"values": copy.deepcopy(event["values"]),
|
| 381 |
+
}
|
| 382 |
+
|
| 383 |
+
|
| 384 |
+
def refresh_event_values(event, tokenizer, prompts):
|
| 385 |
+
event["values"] = {
|
| 386 |
+
field: tokenizer._decode_field(field, tokens, prompts)
|
| 387 |
+
for field, tokens in event["tokens_by_field"].items()
|
| 388 |
+
if tokens
|
| 389 |
+
}
|
| 390 |
+
|
| 391 |
+
|
| 392 |
+
def set_event_local_timestamp(event, tokenizer, local_time):
|
| 393 |
+
time_id = int(round(float(local_time) * tokenizer.time_hz))
|
| 394 |
+
time_id = max(0, min(time_id, tokenizer.n_time_tokens - 1))
|
| 395 |
+
event["tokens_by_field"]["timestamp"] = [tokenizer.time_id_to_token(time_id)]
|
| 396 |
+
|
| 397 |
+
|
| 398 |
+
def active_context_before(events, tokenizer, time_abs):
|
| 399 |
+
state = {}
|
| 400 |
+
for event in events:
|
| 401 |
+
event_time = event.get("time")
|
| 402 |
+
if event_time is None or float(event_time) > float(time_abs) + 1e-6:
|
| 403 |
+
continue
|
| 404 |
+
for field in ("structure", "key", "chord"):
|
| 405 |
+
tokens = event["tokens_by_field"].get(field)
|
| 406 |
+
if tokens:
|
| 407 |
+
state[field] = [int(token) for token in tokens]
|
| 408 |
+
rhythm_tokens = event["tokens_by_field"].get("rhythm", ())
|
| 409 |
+
meter_tokens = [
|
| 410 |
+
int(token)
|
| 411 |
+
for token in rhythm_tokens
|
| 412 |
+
if tokenizer.token_type(token) == "meter"
|
| 413 |
+
]
|
| 414 |
+
if meter_tokens:
|
| 415 |
+
state["meter"] = meter_tokens[:1]
|
| 416 |
+
return state
|
| 417 |
+
|
| 418 |
+
|
| 419 |
+
def apply_prefix_context(event, context, tokenizer, prompts):
|
| 420 |
+
for field in ("structure", "key", "chord"):
|
| 421 |
+
if field not in event["tokens_by_field"] and field in context:
|
| 422 |
+
event["tokens_by_field"][field] = list(context[field])
|
| 423 |
+
|
| 424 |
+
rhythm_tokens = list(event["tokens_by_field"].get("rhythm", ()))
|
| 425 |
+
has_meter = any(tokenizer.token_type(token) == "meter" for token in rhythm_tokens)
|
| 426 |
+
has_eighth = any(
|
| 427 |
+
tokenizer.token_type(token) == "eighth_position" for token in rhythm_tokens
|
| 428 |
+
)
|
| 429 |
+
if has_eighth and not has_meter and "meter" in context:
|
| 430 |
+
event["tokens_by_field"]["rhythm"] = list(context["meter"]) + rhythm_tokens
|
| 431 |
+
|
| 432 |
+
refresh_event_values(event, tokenizer, prompts)
|
| 433 |
+
|
| 434 |
+
|
| 435 |
+
def build_overlap_prefix_tokens(
|
| 436 |
+
stitched_events,
|
| 437 |
+
tokenizer,
|
| 438 |
+
prompts,
|
| 439 |
+
window_start,
|
| 440 |
+
prefix_end,
|
| 441 |
+
):
|
| 442 |
+
eps = 1e-4
|
| 443 |
+
source_events = [
|
| 444 |
+
event
|
| 445 |
+
for event in stitched_events
|
| 446 |
+
if float(window_start) - eps <= float(event.get("time", -1.0)) < float(prefix_end) - eps
|
| 447 |
+
]
|
| 448 |
+
source_events.sort(
|
| 449 |
+
key=lambda event: (
|
| 450 |
+
int(
|
| 451 |
+
event.get(
|
| 452 |
+
"global_subbeat",
|
| 453 |
+
event.get("source_subbeat", event["subbeat"]),
|
| 454 |
+
)
|
| 455 |
+
),
|
| 456 |
+
float(event.get("time", 0.0)),
|
| 457 |
+
)
|
| 458 |
+
)
|
| 459 |
+
first_beat_index = next(
|
| 460 |
+
(
|
| 461 |
+
index
|
| 462 |
+
for index, event in enumerate(source_events)
|
| 463 |
+
if "timestamp" in event["values"] or "rhythm" in event["values"]
|
| 464 |
+
),
|
| 465 |
+
None,
|
| 466 |
+
)
|
| 467 |
+
if first_beat_index is None:
|
| 468 |
+
return None, None, None
|
| 469 |
+
|
| 470 |
+
source_events = source_events[first_beat_index:]
|
| 471 |
+
base_subbeat = int(
|
| 472 |
+
source_events[0].get(
|
| 473 |
+
"global_subbeat",
|
| 474 |
+
source_events[0].get("source_subbeat", source_events[0]["subbeat"]),
|
| 475 |
+
)
|
| 476 |
+
)
|
| 477 |
+
context = active_context_before(
|
| 478 |
+
stitched_events,
|
| 479 |
+
tokenizer,
|
| 480 |
+
source_events[0]["time"],
|
| 481 |
+
)
|
| 482 |
+
prefix_events = []
|
| 483 |
+
for source in source_events:
|
| 484 |
+
event = clone_decoded_event(source)
|
| 485 |
+
source_subbeat = int(
|
| 486 |
+
source.get(
|
| 487 |
+
"global_subbeat",
|
| 488 |
+
source.get("source_subbeat", source["subbeat"]),
|
| 489 |
+
)
|
| 490 |
+
)
|
| 491 |
+
event["subbeat"] = max(0, source_subbeat - base_subbeat)
|
| 492 |
+
if "timestamp" in event["tokens_by_field"]:
|
| 493 |
+
set_event_local_timestamp(
|
| 494 |
+
event,
|
| 495 |
+
tokenizer,
|
| 496 |
+
float(source["time"]) - float(window_start),
|
| 497 |
+
)
|
| 498 |
+
refresh_event_values(event, tokenizer, prompts)
|
| 499 |
+
prefix_events.append(event)
|
| 500 |
+
|
| 501 |
+
apply_prefix_context(prefix_events[0], context, tokenizer, prompts)
|
| 502 |
+
prefix_decoded = {
|
| 503 |
+
"schema_version": tokenizer.schema_version,
|
| 504 |
+
"prompts": prompts,
|
| 505 |
+
"events": prefix_events,
|
| 506 |
+
"has_eos": False,
|
| 507 |
+
}
|
| 508 |
+
return (
|
| 509 |
+
prefix_decoded,
|
| 510 |
+
tokenizer.encode_decoded_sequence(prefix_decoded),
|
| 511 |
+
base_subbeat,
|
| 512 |
+
)
|
| 513 |
+
|
| 514 |
+
|
| 515 |
+
def stitched_window_events(
|
| 516 |
+
decoded,
|
| 517 |
+
time_lookup,
|
| 518 |
+
window_start,
|
| 519 |
+
accept_start,
|
| 520 |
+
accept_end,
|
| 521 |
+
song_duration,
|
| 522 |
+
window_index,
|
| 523 |
+
global_subbeat_base=0,
|
| 524 |
+
):
|
| 525 |
+
accepted = []
|
| 526 |
+
eps = 1e-4
|
| 527 |
+
for event in decoded["events"]:
|
| 528 |
+
local_time = float(time_lookup(event["subbeat"]))
|
| 529 |
+
abs_time = float(window_start) + local_time
|
| 530 |
+
if abs_time < float(accept_start) - eps:
|
| 531 |
+
continue
|
| 532 |
+
if abs_time >= float(accept_end) - eps or abs_time >= float(song_duration) - eps:
|
| 533 |
+
continue
|
| 534 |
+
|
| 535 |
+
output = clone_decoded_event(event)
|
| 536 |
+
output["time"] = float(np.clip(abs_time, 0.0, song_duration))
|
| 537 |
+
output["window_index"] = int(window_index)
|
| 538 |
+
output["window_start"] = float(window_start)
|
| 539 |
+
output["source_subbeat"] = int(event["subbeat"])
|
| 540 |
+
output["global_subbeat"] = int(global_subbeat_base) + int(event["subbeat"])
|
| 541 |
+
if "timestamp" in output["values"]:
|
| 542 |
+
output["values"]["timestamp"] = output["time"]
|
| 543 |
+
|
| 544 |
+
notes = output["values"].get("melody")
|
| 545 |
+
if notes is not None:
|
| 546 |
+
fixed_notes = []
|
| 547 |
+
for note in notes:
|
| 548 |
+
fixed_note = copy.deepcopy(note)
|
| 549 |
+
duration_steps = int(fixed_note["duration_steps"])
|
| 550 |
+
local_end = float(time_lookup(int(event["subbeat"]) + duration_steps))
|
| 551 |
+
end_time = float(window_start) + local_end
|
| 552 |
+
end_time = min(
|
| 553 |
+
float(song_duration),
|
| 554 |
+
max(output["time"] + 0.04, end_time),
|
| 555 |
+
)
|
| 556 |
+
fixed_note["end_time"] = end_time
|
| 557 |
+
fixed_notes.append(fixed_note)
|
| 558 |
+
output["values"]["melody"] = fixed_notes
|
| 559 |
+
accepted.append(output)
|
| 560 |
+
return accepted
|
| 561 |
+
|
| 562 |
+
|
| 563 |
+
def write_window_tokens(path, windows, tokenizer):
|
| 564 |
+
with Path(path).open("w", encoding="utf-8") as f:
|
| 565 |
+
for window in windows:
|
| 566 |
+
f.write(
|
| 567 |
+
f"# window_index={window['window_index']} "
|
| 568 |
+
f"start={window['start']:.6f} end={window['end']:.6f} "
|
| 569 |
+
f"prefix_end={window['prefix_end']:.6f} "
|
| 570 |
+
f"accept=[{window['accept_start']:.6f},{window['accept_end']:.6f}) "
|
| 571 |
+
f"generation_stop={window['generation_stop']} "
|
| 572 |
+
f"prefix_tokens={window['prefix_tokens']} "
|
| 573 |
+
f"tokens={int(window['tokens'].numel())}\n"
|
| 574 |
+
)
|
| 575 |
+
for index, token in enumerate(window["tokens"].tolist()):
|
| 576 |
+
f.write(f"{index}\t{token}\t{tokenizer.describe(token)}\n")
|
| 577 |
+
f.write("\n")
|
| 578 |
+
|
| 579 |
+
|
| 580 |
+
@torch.inference_mode()
|
| 581 |
+
def generate(model, input_values, prompts=None, *, attention_mask=None, max_length=None,
|
| 582 |
+
max_new_tokens=None, prefix_tokens=None,
|
| 583 |
+
autocast_dtype=torch.bfloat16, stop_time_seconds=None,
|
| 584 |
+
return_dict_in_generate=False, output_logits=False, output_scores=False,
|
| 585 |
+
progress_callback=None):
|
| 586 |
+
"""Generate symbolic token sequences with the model's event grammar.
|
| 587 |
+
|
| 588 |
+
Inputs are 24 kHz waveforms [batch, samples]. Scores/logits are per generated
|
| 589 |
+
step, after the prompt prefix, with completed batch items filled by NaNs.
|
| 590 |
+
"""
|
| 591 |
+
from transformers.utils import ModelOutput
|
| 592 |
+
if input_values.ndim == 1:
|
| 593 |
+
input_values = input_values.unsqueeze(0)
|
| 594 |
+
if attention_mask is not None and attention_mask.ndim == 1:
|
| 595 |
+
attention_mask = attention_mask.unsqueeze(0)
|
| 596 |
+
input_values = input_values.to(next(model.parameters()).device)
|
| 597 |
+
prompts = model.tokenizer.normalize_prompts(prompts or FULL_TASK_PROMPTS)
|
| 598 |
+
prefix = list(prefix_tokens) if prefix_tokens is not None else model.tokenizer.prompt_prefix(prompts)
|
| 599 |
+
prefix_length = len(prefix) - (1 if prefix and prefix[-1] == model.tokenizer.eos_token else 0)
|
| 600 |
+
if max_new_tokens is not None:
|
| 601 |
+
if max_length is not None:
|
| 602 |
+
raise ValueError("Specify max_new_tokens or max_length, not both.")
|
| 603 |
+
if not isinstance(max_new_tokens, int) or isinstance(max_new_tokens, bool) or max_new_tokens < 1:
|
| 604 |
+
raise ValueError("max_new_tokens must be a positive integer.")
|
| 605 |
+
max_length = prefix_length + max_new_tokens
|
| 606 |
+
limit = model.max_output_seq_len if max_length is None else max_length
|
| 607 |
+
if not isinstance(limit, int) or isinstance(limit, bool) or not prefix_length < limit <= model.max_output_seq_len:
|
| 608 |
+
raise ValueError("The token limit must exceed the prefix length and fit the decoder context.")
|
| 609 |
+
memory = None
|
| 610 |
+
if attention_mask is not None:
|
| 611 |
+
with inference_autocast(input_values.device, autocast_dtype):
|
| 612 |
+
memory = model.get_audio_features(input_values, attention_mask=attention_mask).encoder_last_hidden_state
|
| 613 |
+
raw, scores, positions = [], [], []
|
| 614 |
+
batch_size = input_values.shape[0]
|
| 615 |
+
def capture(position, ids, logits, masked):
|
| 616 |
+
positions.append(position)
|
| 617 |
+
for enabled, values, output in ((output_logits, logits, raw), (output_scores, masked, scores)):
|
| 618 |
+
if enabled:
|
| 619 |
+
row = torch.full((batch_size, values.shape[-1]), float('nan'), dtype=torch.float32)
|
| 620 |
+
row[ids] = values.detach().cpu()
|
| 621 |
+
output.append(row)
|
| 622 |
+
tokens = constrained_prompt_generate_batch(
|
| 623 |
+
model, input_values, prompts, limit,
|
| 624 |
+
prefix_tokens=prefix_tokens, autocast_dtype=autocast_dtype,
|
| 625 |
+
stop_time_seconds=stop_time_seconds, progress_callback=progress_callback,
|
| 626 |
+
step_callback=capture if output_logits or output_scores else None,
|
| 627 |
+
memory=memory,
|
| 628 |
+
)
|
| 629 |
+
sequences = torch.nn.utils.rnn.pad_sequence(tokens, batch_first=True, padding_value=model.tokenizer.pad_token)
|
| 630 |
+
if return_dict_in_generate or output_logits or output_scores:
|
| 631 |
+
return ModelOutput(sequences=sequences, logits=tuple(raw) if output_logits else None,
|
| 632 |
+
scores=tuple(scores) if output_scores else None,
|
| 633 |
+
token_positions=torch.tensor(positions, dtype=torch.long))
|
| 634 |
+
return sequences
|
infer.py
ADDED
|
@@ -0,0 +1,78 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Transcribe a music recording with SheetSage2."""
|
| 2 |
+
import argparse
|
| 3 |
+
from pathlib import Path
|
| 4 |
+
import sys
|
| 5 |
+
|
| 6 |
+
SCRIPT_DIR = Path(__file__).absolute().parent
|
| 7 |
+
sys.path.insert(0, str(SCRIPT_DIR))
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
def main():
|
| 11 |
+
parser = argparse.ArgumentParser(description="SheetSage2: music audio to ABC, MIDI, annotations and features")
|
| 12 |
+
parser.add_argument("audio", type=Path)
|
| 13 |
+
parser.add_argument("--output", type=Path, default=Path("output"))
|
| 14 |
+
local = SCRIPT_DIR
|
| 15 |
+
parser.add_argument("--model", default=str(local) if (local / "config.json").exists() else "m-a-p/SheetSage2")
|
| 16 |
+
parser.add_argument("--revision", help="Model and code commit on Hugging Face")
|
| 17 |
+
parser.add_argument("--device", default="auto")
|
| 18 |
+
parser.add_argument("--dtype", choices=("bf16", "fp32"), default="bf16")
|
| 19 |
+
parser.add_argument("--preset", choices=("default", "paper"), default="default")
|
| 20 |
+
parser.add_argument("--max-seconds", type=float)
|
| 21 |
+
parser.add_argument("--overlap", type=float)
|
| 22 |
+
parser.add_argument("--lookahead", type=float)
|
| 23 |
+
parser.add_argument("--prompts", nargs="+")
|
| 24 |
+
parser.add_argument("--melody-only", action="store_true",
|
| 25 |
+
help="Keep vocal and instrumental melodies; omit chords from ABC and playback")
|
| 26 |
+
parser.add_argument("--export-logits", action="store_true")
|
| 27 |
+
parser.add_argument("--export-scores", action="store_true")
|
| 28 |
+
parser.add_argument("--export-embeddings", action="store_true")
|
| 29 |
+
parser.add_argument("--all-layers", action="store_true", help="Also export all 24 MERT block features")
|
| 30 |
+
parser.add_argument("--render-audio", action="store_true")
|
| 31 |
+
parser.add_argument("--render-score", nargs="?", const="pdf", default=False, help="pdf,svg,png (default: pdf)")
|
| 32 |
+
parser.add_argument("--render-parts", default="mix", help="mix,melody,vocal,instrumental,chords,all")
|
| 33 |
+
parser.add_argument("--local-files-only", action="store_true")
|
| 34 |
+
args = parser.parse_args()
|
| 35 |
+
if args.render_audio or args.render_score:
|
| 36 |
+
from rendering_sheetsage2 import validate_render_options
|
| 37 |
+
try:
|
| 38 |
+
validate_render_options(audio=args.render_audio, score=args.render_score, parts=args.render_parts)
|
| 39 |
+
except ValueError as exc:
|
| 40 |
+
parser.error(str(exc))
|
| 41 |
+
import torch
|
| 42 |
+
from transformers import AutoModel
|
| 43 |
+
torch.set_num_threads(min(4, torch.get_num_threads()))
|
| 44 |
+
device = "cuda" if torch.cuda.is_available() else "cpu" if args.device == "auto" else args.device
|
| 45 |
+
if args.device != "auto":
|
| 46 |
+
device = args.device
|
| 47 |
+
model = AutoModel.from_pretrained(
|
| 48 |
+
args.model, revision=args.revision, code_revision=args.revision,
|
| 49 |
+
local_files_only=args.local_files_only, trust_remote_code=True,
|
| 50 |
+
).eval().to(device)
|
| 51 |
+
def progress(value):
|
| 52 |
+
if value["stage"] == "encoding":
|
| 53 |
+
print(f"Window {value['window']}/{value['windows']}", flush=True)
|
| 54 |
+
options = dict(dtype=args.dtype, preset=args.preset, max_seconds=args.max_seconds,
|
| 55 |
+
overlap_seconds=args.overlap, lookahead_seconds=args.lookahead,
|
| 56 |
+
export_logits=args.export_logits, export_scores=args.export_scores,
|
| 57 |
+
export_embeddings=args.export_embeddings, output_hidden_states=args.all_layers,
|
| 58 |
+
render_audio=args.render_audio, render_score=args.render_score,
|
| 59 |
+
render_parts=tuple(args.render_parts.split(",")), progress=progress)
|
| 60 |
+
if args.prompts:
|
| 61 |
+
options["prompts"] = args.prompts
|
| 62 |
+
if args.melody_only:
|
| 63 |
+
options["melody_only"] = True
|
| 64 |
+
try:
|
| 65 |
+
result = model.transcribe(args.audio, output_dir=args.output, **options)
|
| 66 |
+
except (ValueError, RuntimeError, FileNotFoundError) as exc:
|
| 67 |
+
parser.exit(1, f"SheetSage2: {exc}\n")
|
| 68 |
+
if args.melody_only and (result.get("abc_error") or not result.get("abc")):
|
| 69 |
+
reason = result.get("abc_error") or "no ABC score was produced"
|
| 70 |
+
parser.exit(1, f"SheetSage2: melody-only ABC unavailable: {reason}. "
|
| 71 |
+
f"Transcription outputs were saved to {args.output.resolve()}.\n")
|
| 72 |
+
print(f"Saved transcription to {args.output.resolve()}")
|
| 73 |
+
if result.get("abc_error"):
|
| 74 |
+
print(f"ABC unavailable: {result['abc_error']}. MIDI and annotations were saved.")
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
if __name__ == "__main__":
|
| 78 |
+
main()
|
io_sheetsage2.py
ADDED
|
@@ -0,0 +1,67 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import os
|
| 2 |
+
import stat
|
| 3 |
+
import tempfile
|
| 4 |
+
from contextlib import contextmanager
|
| 5 |
+
from pathlib import Path
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
def _target_path(path) -> Path:
|
| 9 |
+
target = Path(path)
|
| 10 |
+
target.parent.mkdir(parents=True, exist_ok=True)
|
| 11 |
+
return target
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def _flush_path(path: Path) -> None:
|
| 15 |
+
with path.open("rb+") as handle:
|
| 16 |
+
os.fsync(handle.fileno())
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def _replacement_mode(target: Path) -> int:
|
| 20 |
+
try:
|
| 21 |
+
return stat.S_IMODE(target.stat().st_mode)
|
| 22 |
+
except FileNotFoundError:
|
| 23 |
+
mask = os.umask(0)
|
| 24 |
+
os.umask(mask)
|
| 25 |
+
return 0o666 & ~mask
|
| 26 |
+
|
| 27 |
+
|
| 28 |
+
@contextmanager
|
| 29 |
+
def atomic_output_path(path):
|
| 30 |
+
target = _target_path(path)
|
| 31 |
+
mode = _replacement_mode(target)
|
| 32 |
+
fd, tmp_name = tempfile.mkstemp(
|
| 33 |
+
# Keep the temporary basename independent of the final basename. Some
|
| 34 |
+
# annotation IDs are already close to NAME_MAX; repeating the target
|
| 35 |
+
# name here made an otherwise valid final path fail with ENAMETOOLONG.
|
| 36 |
+
prefix=".tmp_",
|
| 37 |
+
suffix=".tmp",
|
| 38 |
+
dir=str(target.parent),
|
| 39 |
+
)
|
| 40 |
+
os.close(fd)
|
| 41 |
+
tmp_path = Path(tmp_name)
|
| 42 |
+
try:
|
| 43 |
+
yield str(tmp_path)
|
| 44 |
+
os.chmod(tmp_path, mode)
|
| 45 |
+
_flush_path(tmp_path)
|
| 46 |
+
os.replace(tmp_path, target)
|
| 47 |
+
except Exception:
|
| 48 |
+
try:
|
| 49 |
+
tmp_path.unlink()
|
| 50 |
+
except FileNotFoundError:
|
| 51 |
+
pass
|
| 52 |
+
raise
|
| 53 |
+
|
| 54 |
+
|
| 55 |
+
def atomic_write_text(path, text: str, encoding: str = "utf-8", newline: str = "\n") -> str:
|
| 56 |
+
with atomic_output_path(path) as tmp_path:
|
| 57 |
+
with open(tmp_path, "w", encoding=encoding, newline=newline) as handle:
|
| 58 |
+
handle.write(text)
|
| 59 |
+
handle.flush()
|
| 60 |
+
os.fsync(handle.fileno())
|
| 61 |
+
return str(path)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def atomic_write_pretty_midi(midi, path) -> str:
|
| 65 |
+
with atomic_output_path(path) as tmp_path:
|
| 66 |
+
midi.write(tmp_path)
|
| 67 |
+
return str(path)
|
labels_sheetsage2.py
ADDED
|
@@ -0,0 +1,29 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
|
| 2 |
+
|
| 3 |
+
|
| 4 |
+
STRUCTURE_LABELS = ['silence', 'intro', 'outro', 'verse', 'chorus', 'bridge', 'pre-chorus', 'post-chorus',
|
| 5 |
+
'interlude', 'fade-out', 'loop', 'rap', 'preshot', 'irregular']
|
| 6 |
+
|
| 7 |
+
|
| 8 |
+
STRUCTURE_LABELS = ['silence', 'intro', 'outro', 'verse', 'chorus', 'bridge', 'pre-chorus', 'post-chorus',
|
| 9 |
+
'interlude', 'fade-out', 'loop', 'rap', 'preshot', 'irregular', 'instrumental',
|
| 10 |
+
'intro and verse', 'pre-chorus and chorus', 'verse and pre-chorus', 'solo', 'theme', 'development', 'variation', 'pre-outro']
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
MIREX_STRUCTURE_LABEL_REDUCTION = {
|
| 14 |
+
'silence': 'silence',
|
| 15 |
+
'intro': 'intro',
|
| 16 |
+
'outro': 'outro',
|
| 17 |
+
'verse': 'verse',
|
| 18 |
+
'chorus': 'chorus',
|
| 19 |
+
'bridge': 'bridge',
|
| 20 |
+
'pre-chorus': 'verse',
|
| 21 |
+
'post-chorus': 'verse',
|
| 22 |
+
'interlude': 'inst',
|
| 23 |
+
'inst': 'inst',
|
| 24 |
+
'fade-out': 'outro',
|
| 25 |
+
'loop': 'chorus',
|
| 26 |
+
'rap': 'verse',
|
| 27 |
+
'preshot': 'inst',
|
| 28 |
+
'irregular': 'verse',
|
| 29 |
+
}
|
midi_sheetsage2.py
ADDED
|
@@ -0,0 +1,92 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Original-time MIDI playback and a separate score-to-audio measure map."""
|
| 2 |
+
import json
|
| 3 |
+
from io import BytesIO
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
|
| 6 |
+
import mir_eval.chord
|
| 7 |
+
import numpy as np
|
| 8 |
+
import pretty_midi
|
| 9 |
+
|
| 10 |
+
from .io_sheetsage2 import atomic_write_text
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def chord_pitches(label):
|
| 14 |
+
if label in {"N", "X", "?"}:
|
| 15 |
+
return []
|
| 16 |
+
root, bitmap, bass = mir_eval.chord.encode(label, reduce_extended_chords=True)
|
| 17 |
+
if root < 0:
|
| 18 |
+
return []
|
| 19 |
+
# Keep inversion bass below the chord; no ABC chord-name reinterpretation.
|
| 20 |
+
upper = [48 + root + int(interval) for interval in np.flatnonzero(bitmap > 0)]
|
| 21 |
+
return sorted(set([36 + (root + bass) % 12] + upper))
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def measure_map(score, duration):
|
| 25 |
+
rows, position = [], 0.0
|
| 26 |
+
for measure in score.measures:
|
| 27 |
+
length = measure.abc_numerator / measure.abc_denominator
|
| 28 |
+
actual = measure.numerator / measure.denominator
|
| 29 |
+
start = float(score.beats[measure.start_beat].time)
|
| 30 |
+
end = float(score.beats[measure.end_beat].time)
|
| 31 |
+
rows.append(dict(index=measure.index, start=start, end=min(end, duration),
|
| 32 |
+
score_start=position, score_end=position + length,
|
| 33 |
+
leading_rest=length - actual if measure.pad_before else 0.0,
|
| 34 |
+
trailing_rest=length - actual if not measure.pad_before else 0.0))
|
| 35 |
+
position += length
|
| 36 |
+
return rows
|
| 37 |
+
|
| 38 |
+
|
| 39 |
+
def midi_bytes(midi):
|
| 40 |
+
"""Serialize a PrettyMIDI object without touching the filesystem."""
|
| 41 |
+
stream = BytesIO()
|
| 42 |
+
midi.write(stream)
|
| 43 |
+
return stream.getvalue()
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def build_playback(melody_midi, chord_rows, score, duration):
|
| 47 |
+
"""Return playback metadata and MIDI bytes, preserving MIDI tick rounding."""
|
| 48 |
+
if isinstance(melody_midi, pretty_midi.PrettyMIDI):
|
| 49 |
+
melody_midi = midi_bytes(melody_midi)
|
| 50 |
+
midi = pretty_midi.PrettyMIDI(BytesIO(melody_midi))
|
| 51 |
+
chord_track = pretty_midi.Instrument(0, name="Chords")
|
| 52 |
+
warnings = []
|
| 53 |
+
for start, end, label in chord_rows:
|
| 54 |
+
start, end = max(0.0, float(start)), min(duration, float(end))
|
| 55 |
+
try:
|
| 56 |
+
pitches = chord_pitches(label)
|
| 57 |
+
except mir_eval.chord.InvalidChordException as exc:
|
| 58 |
+
warnings.append(f"Chord playback skipped {label}: {exc}")
|
| 59 |
+
continue
|
| 60 |
+
# Rearticulate long chord spans at downbeats, including repeated bars.
|
| 61 |
+
downbeats = [b.time for b in score.beats if b.beat_id == 1] if score is not None else []
|
| 62 |
+
cuts = [start] + [time for time in downbeats if start < time < end] + [end]
|
| 63 |
+
for a, b in zip(cuts, cuts[1:]):
|
| 64 |
+
if b <= a:
|
| 65 |
+
continue
|
| 66 |
+
for pitch in pitches:
|
| 67 |
+
chord_track.notes.append(pretty_midi.Note(48, pitch, a, b))
|
| 68 |
+
midi.instruments.append(chord_track)
|
| 69 |
+
transcription = midi_bytes(midi)
|
| 70 |
+
chords = pretty_midi.PrettyMIDI(resolution=midi.resolution)
|
| 71 |
+
chords.instruments.append(chord_track)
|
| 72 |
+
midis = {"transcription": transcription, "chords": midi_bytes(chords)}
|
| 73 |
+
# Decode the serialized bytes, so playback follows the returned MIDI exactly.
|
| 74 |
+
midi = pretty_midi.PrettyMIDI(BytesIO(transcription))
|
| 75 |
+
tracks = [dict(name=instrument.name, program=int(instrument.program),
|
| 76 |
+
notes=[dict(pitch=int(n.pitch), start=float(n.start), end=float(n.end), velocity=int(n.velocity))
|
| 77 |
+
for n in instrument.notes]) for instrument in midi.instruments]
|
| 78 |
+
data = dict(version=1, duration=duration, midi="transcription.mid", tracks=tracks,
|
| 79 |
+
measures=measure_map(score, duration) if score is not None else [], warnings=warnings)
|
| 80 |
+
return data, midis
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
def export_playback(directory, score, duration):
|
| 84 |
+
"""Write playback files for an existing directory of exported annotations."""
|
| 85 |
+
directory = Path(directory)
|
| 86 |
+
chord_path = directory / "chord.lab"
|
| 87 |
+
rows = [line.split(maxsplit=2) for line in chord_path.read_text(encoding="utf-8").splitlines()] if chord_path.exists() else []
|
| 88 |
+
data, midis = build_playback((directory / "melody.mid").read_bytes(), rows, score, duration)
|
| 89 |
+
for name, content in midis.items():
|
| 90 |
+
(directory / f"{name}.mid").write_bytes(content)
|
| 91 |
+
atomic_write_text(directory / "playback.json", json.dumps(data, indent=2))
|
| 92 |
+
return data
|
model.safetensors
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:b235f68091a5f5b644000f2b5acb57d1e70432aca2b34ab1b9cf27236e1f4274
|
| 3 |
+
size 228738564
|
modeling_mert2.py
ADDED
|
@@ -0,0 +1,361 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Standalone MERT2 waveform-to-representation inference."""
|
| 2 |
+
|
| 3 |
+
from dataclasses import dataclass
|
| 4 |
+
from typing import Optional, Tuple
|
| 5 |
+
|
| 6 |
+
import torch
|
| 7 |
+
from torch import nn
|
| 8 |
+
from torch.nn import functional as F
|
| 9 |
+
from torchaudio.transforms import AmplitudeToDB, MelScale, Spectrogram
|
| 10 |
+
from transformers import PreTrainedModel
|
| 11 |
+
from transformers.utils import ModelOutput
|
| 12 |
+
|
| 13 |
+
from .configuration_mert2 import MERT2Config
|
| 14 |
+
|
| 15 |
+
|
| 16 |
+
@dataclass
|
| 17 |
+
class MERT2ModelOutput(ModelOutput):
|
| 18 |
+
"""Frame representations and their validity mask.
|
| 19 |
+
|
| 20 |
+
``hidden_states`` contains one tensor per Conformer block, without an input
|
| 21 |
+
embedding. ``feature_attention_mask`` is True at valid output frames.
|
| 22 |
+
"""
|
| 23 |
+
|
| 24 |
+
last_hidden_state: Optional[torch.FloatTensor] = None
|
| 25 |
+
hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
|
| 26 |
+
feature_attention_mask: Optional[torch.BoolTensor] = None
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class MERT2MelFrontend(nn.Module):
|
| 30 |
+
"""Power log-mel features with fixed, checkpoint-specific normalization."""
|
| 31 |
+
|
| 32 |
+
def __init__(self, config):
|
| 33 |
+
super().__init__()
|
| 34 |
+
# Filter construction requires real tensors even during meta loading.
|
| 35 |
+
with torch.device("cpu"):
|
| 36 |
+
self.register_buffer("mel_mean", torch.zeros(config.num_mel_bins, dtype=torch.float32))
|
| 37 |
+
self.register_buffer("mel_std", torch.ones(config.num_mel_bins, dtype=torch.float32))
|
| 38 |
+
self.spectrogram = Spectrogram(
|
| 39 |
+
n_fft=config.n_fft,
|
| 40 |
+
win_length=config.win_length,
|
| 41 |
+
hop_length=config.hop_length,
|
| 42 |
+
power=2.0,
|
| 43 |
+
).float()
|
| 44 |
+
self.mel_scale = MelScale(
|
| 45 |
+
n_mels=config.num_mel_bins,
|
| 46 |
+
sample_rate=config.sampling_rate,
|
| 47 |
+
n_stft=config.n_fft // 2 + 1,
|
| 48 |
+
).float()
|
| 49 |
+
self.amplitude_to_db = AmplitudeToDB(stype="power", top_db=None)
|
| 50 |
+
|
| 51 |
+
def _apply(self, fn, recurse=True):
|
| 52 |
+
# Preserve the original float32 values, including on model.half().
|
| 53 |
+
saved = [(module, name, value) for module in self.modules() for name, value in module._buffers.items() if value is not None]
|
| 54 |
+
super()._apply(fn, recurse=recurse)
|
| 55 |
+
for module, name, original in saved:
|
| 56 |
+
current = module._buffers[name]
|
| 57 |
+
if original.is_floating_point():
|
| 58 |
+
module._buffers[name] = (
|
| 59 |
+
current.float() if original.is_meta else original.to(device=current.device, dtype=torch.float32)
|
| 60 |
+
)
|
| 61 |
+
return self
|
| 62 |
+
|
| 63 |
+
@torch.no_grad()
|
| 64 |
+
def forward(self, waveform):
|
| 65 |
+
with torch.autocast(device_type=waveform.device.type, enabled=False):
|
| 66 |
+
spectrum = self.spectrogram(waveform.float())
|
| 67 |
+
mel = self.amplitude_to_db(self.mel_scale(spectrum))
|
| 68 |
+
mel = mel[..., :-1].transpose(-1, -2)
|
| 69 |
+
return (mel - self.mel_mean) / self.mel_std.clamp_min(1e-5)
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
class Transpose(nn.Module):
|
| 73 |
+
def forward(self, hidden_states):
|
| 74 |
+
return hidden_states.transpose(1, 2)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
class GlobalResponseNorm(nn.Module):
|
| 78 |
+
def __init__(self, dim):
|
| 79 |
+
super().__init__()
|
| 80 |
+
self.weight = nn.Parameter(torch.zeros(1, 1, dim))
|
| 81 |
+
self.bias = nn.Parameter(torch.zeros(1, 1, dim))
|
| 82 |
+
|
| 83 |
+
def forward(self, hidden_states):
|
| 84 |
+
magnitude = torch.norm(hidden_states, p=2, dim=1, keepdim=True)
|
| 85 |
+
normalized = magnitude / (magnitude.mean(dim=-1, keepdim=True) + 1e-6)
|
| 86 |
+
return self.weight * (hidden_states * normalized) + self.bias + hidden_states
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
class ConvNextLayer(nn.Module):
|
| 90 |
+
def __init__(self, dim, eps):
|
| 91 |
+
super().__init__()
|
| 92 |
+
self.depthwise_block = nn.Sequential(
|
| 93 |
+
Transpose(), nn.Conv1d(dim, dim, 7, padding=3, groups=dim), Transpose()
|
| 94 |
+
)
|
| 95 |
+
self.pointwise_block = nn.Sequential(
|
| 96 |
+
nn.LayerNorm(dim, eps=eps),
|
| 97 |
+
nn.Linear(dim, 4 * dim),
|
| 98 |
+
nn.GELU(),
|
| 99 |
+
GlobalResponseNorm(4 * dim),
|
| 100 |
+
nn.Linear(4 * dim, dim),
|
| 101 |
+
)
|
| 102 |
+
|
| 103 |
+
def forward(self, hidden_states):
|
| 104 |
+
return hidden_states + self.pointwise_block(self.depthwise_block(hidden_states))
|
| 105 |
+
|
| 106 |
+
|
| 107 |
+
class ConvNextBlock(nn.Module):
|
| 108 |
+
def __init__(self, in_channels, out_channels, stride, depth, eps):
|
| 109 |
+
super().__init__()
|
| 110 |
+
self.resampling_layer = (
|
| 111 |
+
nn.Sequential(
|
| 112 |
+
nn.LayerNorm(in_channels, eps=eps),
|
| 113 |
+
Transpose(),
|
| 114 |
+
nn.Conv1d(in_channels, out_channels, 2, stride=stride),
|
| 115 |
+
Transpose(),
|
| 116 |
+
)
|
| 117 |
+
if in_channels != out_channels or stride > 1 else nn.Identity()
|
| 118 |
+
)
|
| 119 |
+
self.convnext_layers = nn.Sequential(*[ConvNextLayer(out_channels, eps) for _ in range(depth)])
|
| 120 |
+
|
| 121 |
+
def forward(self, hidden_states):
|
| 122 |
+
return self.convnext_layers(self.resampling_layer(hidden_states))
|
| 123 |
+
|
| 124 |
+
|
| 125 |
+
class RotaryEmbedding(nn.Module):
|
| 126 |
+
def __init__(self, config):
|
| 127 |
+
super().__init__()
|
| 128 |
+
self.head_dim = config.hidden_size // config.num_attention_heads
|
| 129 |
+
self.base = config.rotary_embedding_base
|
| 130 |
+
with torch.device("cpu"):
|
| 131 |
+
inverse_frequency = 1.0 / (
|
| 132 |
+
self.base ** (torch.arange(0, self.head_dim, 2, dtype=torch.float32) / self.head_dim)
|
| 133 |
+
)
|
| 134 |
+
self.register_buffer("inv_freq", inverse_frequency, persistent=False)
|
| 135 |
+
self._sequence_length = 0
|
| 136 |
+
self._cache_device = None
|
| 137 |
+
self._cos = None
|
| 138 |
+
self._sin = None
|
| 139 |
+
|
| 140 |
+
def _apply(self, fn, recurse=True):
|
| 141 |
+
original = self.inv_freq
|
| 142 |
+
super()._apply(fn, recurse=recurse)
|
| 143 |
+
self.inv_freq = original.to(device=self.inv_freq.device, dtype=torch.float32)
|
| 144 |
+
return self
|
| 145 |
+
|
| 146 |
+
def forward(self, hidden_states):
|
| 147 |
+
length = hidden_states.shape[1]
|
| 148 |
+
if self._cos is None or self._cache_device != hidden_states.device or length > self._sequence_length:
|
| 149 |
+
positions = torch.arange(length, device=hidden_states.device, dtype=self.inv_freq.dtype)
|
| 150 |
+
frequencies = torch.einsum("i,j->ij", positions, self.inv_freq)
|
| 151 |
+
angles = torch.cat((frequencies, frequencies), dim=-1)
|
| 152 |
+
self._cos = angles.cos()[:, None, None, :]
|
| 153 |
+
self._sin = angles.sin()[:, None, None, :]
|
| 154 |
+
self._sequence_length = length
|
| 155 |
+
self._cache_device = hidden_states.device
|
| 156 |
+
return (
|
| 157 |
+
self._cos[:length].to(hidden_states.dtype).permute(1, 0, 2, 3),
|
| 158 |
+
self._sin[:length].to(hidden_states.dtype).permute(1, 0, 2, 3),
|
| 159 |
+
)
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
def rotate_half(value):
|
| 163 |
+
first, second = value.chunk(2, dim=-1)
|
| 164 |
+
return torch.cat((-second, first), dim=-1)
|
| 165 |
+
|
| 166 |
+
|
| 167 |
+
class SelfAttention(nn.Module):
|
| 168 |
+
def __init__(self, config):
|
| 169 |
+
super().__init__()
|
| 170 |
+
self.config = config
|
| 171 |
+
self.num_heads = config.num_attention_heads
|
| 172 |
+
self.head_dim = config.hidden_size // self.num_heads
|
| 173 |
+
self.query_proj = nn.Linear(config.hidden_size, config.hidden_size)
|
| 174 |
+
self.key_proj = nn.Linear(config.hidden_size, config.hidden_size)
|
| 175 |
+
self.value_proj = nn.Linear(config.hidden_size, config.hidden_size)
|
| 176 |
+
self.out_proj = nn.Linear(config.hidden_size, config.hidden_size)
|
| 177 |
+
|
| 178 |
+
def forward(self, hidden_states, position_embeddings):
|
| 179 |
+
batch, time, width = hidden_states.shape
|
| 180 |
+
shape = (batch, time, self.num_heads, self.head_dim)
|
| 181 |
+
query = self.query_proj(hidden_states).reshape(shape)
|
| 182 |
+
key = self.key_proj(hidden_states).reshape(shape)
|
| 183 |
+
value = self.value_proj(hidden_states).reshape(shape)
|
| 184 |
+
cos, sin = position_embeddings
|
| 185 |
+
query = query * cos + rotate_half(query) * sin
|
| 186 |
+
key = key * cos + rotate_half(key) * sin
|
| 187 |
+
if self.config._attn_implementation == "flash_attention_2":
|
| 188 |
+
if query.device.type != "cuda" or query.dtype not in (torch.float16, torch.bfloat16):
|
| 189 |
+
raise ValueError("flash_attention_2 requires CUDA and float16/bfloat16 activations; use autocast or load a reduced-precision model.")
|
| 190 |
+
try:
|
| 191 |
+
from flash_attn import flash_attn_func
|
| 192 |
+
except ImportError as error:
|
| 193 |
+
raise ImportError("Install flash-attn to use flash_attention_2, or select attn_implementation='sdpa'.") from error
|
| 194 |
+
attended = flash_attn_func(query.contiguous(), key.contiguous(), value.contiguous(), dropout_p=0.0, causal=False)
|
| 195 |
+
else:
|
| 196 |
+
attended = F.scaled_dot_product_attention(
|
| 197 |
+
query.transpose(1, 2), key.transpose(1, 2), value.transpose(1, 2), dropout_p=0.0, is_causal=False
|
| 198 |
+
).transpose(1, 2)
|
| 199 |
+
return self.out_proj(attended.reshape(batch, time, width))
|
| 200 |
+
|
| 201 |
+
|
| 202 |
+
class FeedForward(nn.Module):
|
| 203 |
+
def __init__(self, config):
|
| 204 |
+
super().__init__()
|
| 205 |
+
self.w_1 = nn.Linear(config.hidden_size, config.intermediate_size)
|
| 206 |
+
self.w_2 = nn.Linear(config.intermediate_size, config.hidden_size)
|
| 207 |
+
|
| 208 |
+
def forward(self, hidden_states):
|
| 209 |
+
return self.w_2(F.gelu(self.w_1(hidden_states)))
|
| 210 |
+
|
| 211 |
+
|
| 212 |
+
class ConvolutionModule(nn.Module):
|
| 213 |
+
def __init__(self, config):
|
| 214 |
+
super().__init__()
|
| 215 |
+
width = config.hidden_size
|
| 216 |
+
kernel = config.conv_depthwise_kernel_size
|
| 217 |
+
self.layer_norm = nn.LayerNorm(width, eps=config.layer_norm_eps)
|
| 218 |
+
self.conv_block = nn.Sequential(
|
| 219 |
+
Transpose(),
|
| 220 |
+
nn.Conv1d(width, 2 * width, 1, bias=False),
|
| 221 |
+
nn.GLU(dim=1),
|
| 222 |
+
nn.Conv1d(width, width, kernel, padding=(kernel - 1) // 2, groups=width, bias=False),
|
| 223 |
+
nn.Sequential(Transpose(), nn.LayerNorm(width, eps=config.layer_norm_eps), Transpose()),
|
| 224 |
+
nn.GELU(),
|
| 225 |
+
nn.Conv1d(width, width, 1, bias=False),
|
| 226 |
+
Transpose(),
|
| 227 |
+
)
|
| 228 |
+
|
| 229 |
+
def forward(self, hidden_states):
|
| 230 |
+
return self.conv_block(self.layer_norm(hidden_states))
|
| 231 |
+
|
| 232 |
+
|
| 233 |
+
class ConformerBlock(nn.Module):
|
| 234 |
+
def __init__(self, config):
|
| 235 |
+
super().__init__()
|
| 236 |
+
width, eps = config.hidden_size, config.layer_norm_eps
|
| 237 |
+
self.ffn1_layer_norm = nn.LayerNorm(width, eps=eps)
|
| 238 |
+
self.ffn1 = FeedForward(config)
|
| 239 |
+
self.attn_layer_norm = nn.LayerNorm(width, eps=eps)
|
| 240 |
+
self.attn = SelfAttention(config)
|
| 241 |
+
self.conv_module = ConvolutionModule(config)
|
| 242 |
+
self.ffn2_layer_norm = nn.LayerNorm(width, eps=eps)
|
| 243 |
+
self.ffn2 = FeedForward(config)
|
| 244 |
+
self.final_layer_norm = nn.LayerNorm(width, eps=eps)
|
| 245 |
+
|
| 246 |
+
def forward(self, hidden_states, position_embeddings):
|
| 247 |
+
hidden_states = hidden_states + 0.5 * self.ffn1(self.ffn1_layer_norm(hidden_states))
|
| 248 |
+
hidden_states = self.attn(self.attn_layer_norm(hidden_states), position_embeddings) + hidden_states
|
| 249 |
+
hidden_states = self.conv_module(hidden_states) + hidden_states
|
| 250 |
+
hidden_states = hidden_states + 0.5 * self.ffn2(self.ffn2_layer_norm(hidden_states))
|
| 251 |
+
return self.final_layer_norm(hidden_states)
|
| 252 |
+
|
| 253 |
+
|
| 254 |
+
class MERT2Model(PreTrainedModel):
|
| 255 |
+
"""Encode mono waveforms into 25 Hz MERT2 frame representations.
|
| 256 |
+
|
| 257 |
+
Inputs are floating-point mono waveforms sampled at ``config.sampling_rate``.
|
| 258 |
+
Do not standardize their amplitude. An optional waveform attention mask must
|
| 259 |
+
contain a prefix of ones followed by zeros. Without a mask, every input
|
| 260 |
+
sample, including any caller-provided silence or padding, is processed.
|
| 261 |
+
"""
|
| 262 |
+
|
| 263 |
+
config_class = MERT2Config
|
| 264 |
+
base_model_prefix = ""
|
| 265 |
+
main_input_name = "input_values"
|
| 266 |
+
_supports_sdpa = True
|
| 267 |
+
_supports_flash_attn_2 = True
|
| 268 |
+
_no_split_modules = ["ConvNextBlock", "ConformerBlock"]
|
| 269 |
+
|
| 270 |
+
def __init__(self, config):
|
| 271 |
+
super().__init__(config)
|
| 272 |
+
if config._attn_implementation not in {"sdpa", "flash_attention_2"}:
|
| 273 |
+
raise ValueError("MERT2 supports attn_implementation='sdpa' or 'flash_attention_2'.")
|
| 274 |
+
self.feature_extractor = MERT2MelFrontend(config)
|
| 275 |
+
channels = [config.num_mel_bins] + config.subsampling_channels
|
| 276 |
+
self.subsampling_module = nn.Sequential(*[
|
| 277 |
+
ConvNextBlock(channels[i], channels[i + 1], (1, 2, 2)[i], config.subsampling_depths[i], config.subsampling_layer_norm_eps)
|
| 278 |
+
for i in range(3)
|
| 279 |
+
])
|
| 280 |
+
self.layers = nn.ModuleList([ConformerBlock(config) for _ in range(config.num_hidden_layers)])
|
| 281 |
+
self.embed_positions = RotaryEmbedding(config)
|
| 282 |
+
self.post_init()
|
| 283 |
+
|
| 284 |
+
def _init_weights(self, module):
|
| 285 |
+
if isinstance(module, (nn.Linear, nn.Conv1d)):
|
| 286 |
+
nn.init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
|
| 287 |
+
if module.bias is not None:
|
| 288 |
+
nn.init.zeros_(module.bias)
|
| 289 |
+
elif isinstance(module, nn.LayerNorm):
|
| 290 |
+
nn.init.ones_(module.weight)
|
| 291 |
+
nn.init.zeros_(module.bias)
|
| 292 |
+
|
| 293 |
+
def _get_feat_extract_output_lengths(self, input_lengths):
|
| 294 |
+
return input_lengths // self.config.inputs_to_logits_ratio
|
| 295 |
+
|
| 296 |
+
def _encode(self, input_values, output_hidden_states):
|
| 297 |
+
mel = self.feature_extractor(input_values)
|
| 298 |
+
input_dtype = self.subsampling_module[0].convnext_layers[0].depthwise_block[1].weight.dtype
|
| 299 |
+
hidden = self.subsampling_module(mel.to(dtype=input_dtype))
|
| 300 |
+
positions = self.embed_positions(hidden)
|
| 301 |
+
states = [] if output_hidden_states else None
|
| 302 |
+
for layer in self.layers:
|
| 303 |
+
hidden = layer(hidden, positions)
|
| 304 |
+
if states is not None:
|
| 305 |
+
states.append(hidden)
|
| 306 |
+
return hidden, tuple(states) if states is not None else None
|
| 307 |
+
|
| 308 |
+
def forward(self, input_values, attention_mask=None, output_hidden_states=None, return_dict=None):
|
| 309 |
+
"""Return final frames, optionally all block states, and a frame mask.
|
| 310 |
+
|
| 311 |
+
Different valid waveform lengths are encoded separately so convolution
|
| 312 |
+
and global normalization never incorporate another sample's padding.
|
| 313 |
+
Outputs are zero-padded to the longest valid feature sequence.
|
| 314 |
+
"""
|
| 315 |
+
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
|
| 316 |
+
return_dict = self.config.use_return_dict if return_dict is None else return_dict
|
| 317 |
+
if input_values.ndim != 2 or not input_values.is_floating_point() or input_values.shape[0] == 0:
|
| 318 |
+
raise ValueError("input_values must be a nonempty floating-point tensor of shape [batch, samples].")
|
| 319 |
+
if input_values.shape[1] < self.config.minimum_input_samples:
|
| 320 |
+
raise ValueError(f"Each waveform must contain at least {self.config.minimum_input_samples} samples.")
|
| 321 |
+
if not torch.isfinite(input_values).all():
|
| 322 |
+
raise ValueError("input_values must contain only finite samples.")
|
| 323 |
+
batch, samples = input_values.shape
|
| 324 |
+
if attention_mask is None:
|
| 325 |
+
hidden, states = self._encode(input_values, output_hidden_states)
|
| 326 |
+
feature_mask = torch.ones(hidden.shape[:2], dtype=torch.bool, device=hidden.device)
|
| 327 |
+
else:
|
| 328 |
+
if attention_mask.shape != input_values.shape:
|
| 329 |
+
raise ValueError("attention_mask must have the same shape as input_values.")
|
| 330 |
+
attention_mask = attention_mask.to(device=input_values.device)
|
| 331 |
+
if not ((attention_mask == 0) | (attention_mask == 1)).all():
|
| 332 |
+
raise ValueError("attention_mask must contain only zeros and ones.")
|
| 333 |
+
mask = attention_mask.bool()
|
| 334 |
+
lengths = mask.sum(dim=1)
|
| 335 |
+
expected = torch.arange(samples, device=input_values.device)[None, :] < lengths[:, None]
|
| 336 |
+
if not torch.equal(mask, expected):
|
| 337 |
+
raise ValueError("attention_mask must be right padded: valid samples followed by padding.")
|
| 338 |
+
if (lengths < self.config.minimum_input_samples).any():
|
| 339 |
+
raise ValueError(f"Each waveform must contain at least {self.config.minimum_input_samples} valid samples.")
|
| 340 |
+
feature_lengths = self._get_feat_extract_output_lengths(lengths)
|
| 341 |
+
maximum = int(feature_lengths.max().item())
|
| 342 |
+
feature_mask = torch.arange(maximum, device=input_values.device)[None, :] < feature_lengths[:, None]
|
| 343 |
+
hidden, state_values = None, None
|
| 344 |
+
for length in torch.unique(lengths, sorted=True).tolist():
|
| 345 |
+
indices = torch.where(lengths == length)[0]
|
| 346 |
+
group_hidden, group_states = self._encode(input_values.index_select(0, indices)[:, :length], output_hidden_states)
|
| 347 |
+
padding = (0, 0, 0, maximum - group_hidden.shape[1])
|
| 348 |
+
if hidden is None:
|
| 349 |
+
hidden = group_hidden.new_zeros(batch, maximum, self.config.hidden_size)
|
| 350 |
+
if output_hidden_states:
|
| 351 |
+
state_values = [torch.zeros_like(hidden) for _ in self.layers]
|
| 352 |
+
hidden = hidden.index_copy(0, indices, F.pad(group_hidden, padding))
|
| 353 |
+
if state_values is not None:
|
| 354 |
+
for i, value in enumerate(group_states):
|
| 355 |
+
state_values[i] = state_values[i].index_copy(0, indices, F.pad(value, padding))
|
| 356 |
+
states = tuple(state_values) if state_values is not None else None
|
| 357 |
+
output = MERT2ModelOutput(last_hidden_state=hidden, hidden_states=states, feature_attention_mask=feature_mask)
|
| 358 |
+
return output if return_dict else output.to_tuple()
|
| 359 |
+
|
| 360 |
+
|
| 361 |
+
MERT2Model.register_for_auto_class("AutoModel")
|
modeling_sheetsage2.py
ADDED
|
@@ -0,0 +1,448 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Hugging Face SheetSage2 model with a shared MERT-v2 encoder."""
|
| 2 |
+
|
| 3 |
+
from dataclasses import dataclass
|
| 4 |
+
import copy
|
| 5 |
+
import hashlib
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
import shutil
|
| 8 |
+
from types import SimpleNamespace
|
| 9 |
+
from typing import Optional, Tuple
|
| 10 |
+
|
| 11 |
+
import torch
|
| 12 |
+
from torch import nn
|
| 13 |
+
from torch.nn import functional as F
|
| 14 |
+
from transformers import AutoModel, BartConfig, PreTrainedModel
|
| 15 |
+
from transformers.models.bart.modeling_bart import BartDecoder
|
| 16 |
+
from transformers.utils import ModelOutput
|
| 17 |
+
from transformers.utils.hub import cached_file
|
| 18 |
+
|
| 19 |
+
from .configuration_mert2 import MERT2Config
|
| 20 |
+
from .modeling_mert2 import MERT2Model
|
| 21 |
+
from .configuration_sheetsage2 import SheetSage2Config
|
| 22 |
+
from .tokenization_sheetsage2 import SheetSage2Tokenizer
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
PROJECTIONS = ("query_proj", "key_proj", "value_proj", "out_proj")
|
| 26 |
+
BASE_CODE_HASHES = {
|
| 27 |
+
"configuration_mert2.py": "77b53ec9d7ee31a599d744fb006e812c7eeaf7390deb46e2f460cf8c17b00bd6",
|
| 28 |
+
"modeling_mert2.py": "b1a3174e5649c4b26b0c90d8626f0adacfbbba111a58ed3bb72ad651945a2f5c",
|
| 29 |
+
}
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def _hf_relative_dependencies():
|
| 33 |
+
# Transformers 4.45 copies direct imports into a fresh local module cache.
|
| 34 |
+
from .audio_sheetsage2 import load_audio
|
| 35 |
+
from .durations_sheetsage2 import DURATION_TEMPLATES
|
| 36 |
+
from .exports_sheetsage2 import export_result
|
| 37 |
+
from .io_sheetsage2 import atomic_write_text
|
| 38 |
+
from .labels_sheetsage2 import STRUCTURE_LABELS
|
| 39 |
+
from .midi_sheetsage2 import export_playback
|
| 40 |
+
from .notation_sheetsage2 import generate_abc_from_exports
|
| 41 |
+
from .rendering_sheetsage2 import render_outputs
|
| 42 |
+
from .schema_sheetsage2 import get_prompt_multitask_schema
|
| 43 |
+
from .tensors_sheetsage2 import WindowTensorWriter
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
def _sha256(path):
|
| 47 |
+
digest = hashlib.sha256()
|
| 48 |
+
with Path(path).open("rb") as stream:
|
| 49 |
+
for block in iter(lambda: stream.read(8 << 20), b""):
|
| 50 |
+
digest.update(block)
|
| 51 |
+
return digest.hexdigest()
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
@dataclass
|
| 55 |
+
class SheetSage2EncoderOutput(ModelOutput):
|
| 56 |
+
"""Frame features. Block states exclude the separately returned input state."""
|
| 57 |
+
|
| 58 |
+
encoder_last_hidden_state: Optional[torch.FloatTensor] = None
|
| 59 |
+
backbone_last_hidden_state: Optional[torch.FloatTensor] = None
|
| 60 |
+
mixed_hidden_state: Optional[torch.FloatTensor] = None
|
| 61 |
+
input_hidden_state: Optional[torch.FloatTensor] = None
|
| 62 |
+
backbone_hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
|
| 63 |
+
feature_attention_mask: Optional[torch.BoolTensor] = None
|
| 64 |
+
|
| 65 |
+
@property
|
| 66 |
+
def last_hidden_state(self):
|
| 67 |
+
return self.encoder_last_hidden_state
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
@dataclass
|
| 71 |
+
class SheetSage2Output(ModelOutput):
|
| 72 |
+
logits: Optional[torch.FloatTensor] = None
|
| 73 |
+
past_key_values: Optional[Tuple] = None
|
| 74 |
+
decoder_hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
|
| 75 |
+
encoder_last_hidden_state: Optional[torch.FloatTensor] = None
|
| 76 |
+
backbone_last_hidden_state: Optional[torch.FloatTensor] = None
|
| 77 |
+
mixed_hidden_state: Optional[torch.FloatTensor] = None
|
| 78 |
+
input_hidden_state: Optional[torch.FloatTensor] = None
|
| 79 |
+
backbone_hidden_states: Optional[Tuple[torch.FloatTensor, ...]] = None
|
| 80 |
+
feature_attention_mask: Optional[torch.BoolTensor] = None
|
| 81 |
+
|
| 82 |
+
|
| 83 |
+
class AttentionAdapter(nn.Module):
|
| 84 |
+
def __init__(self, width, rank):
|
| 85 |
+
super().__init__()
|
| 86 |
+
self.lora_A = nn.Linear(width, rank, bias=False)
|
| 87 |
+
self.lora_B = nn.Linear(rank, width, bias=False)
|
| 88 |
+
|
| 89 |
+
|
| 90 |
+
class EncoderAdapters(nn.Module):
|
| 91 |
+
def __init__(self, config):
|
| 92 |
+
super().__init__()
|
| 93 |
+
self.layers = nn.ModuleList()
|
| 94 |
+
for _ in range(config.backbone_config["num_hidden_layers"]):
|
| 95 |
+
layer = nn.Module()
|
| 96 |
+
layer.attn = nn.Module()
|
| 97 |
+
for name in PROJECTIONS:
|
| 98 |
+
setattr(layer.attn, name, AttentionAdapter(config.backbone_config["hidden_size"], config.lora_rank))
|
| 99 |
+
self.layers.append(layer)
|
| 100 |
+
|
| 101 |
+
|
| 102 |
+
class SheetSage2Model(PreTrainedModel):
|
| 103 |
+
"""Audio-to-symbolic model, with optional frame and decoder representations.
|
| 104 |
+
|
| 105 |
+
``from_pretrained`` loads the pinned MERT-v2 parent and merges attention
|
| 106 |
+
adapters in float32. ``save_pretrained`` writes an independent merged model.
|
| 107 |
+
``forward`` returns raw vocabulary logits; grammar-masked generation scores
|
| 108 |
+
are available separately through ``generate(output_scores=True)``.
|
| 109 |
+
"""
|
| 110 |
+
|
| 111 |
+
config_class = SheetSage2Config
|
| 112 |
+
base_model_prefix = ""
|
| 113 |
+
main_input_name = "input_values"
|
| 114 |
+
_supports_sdpa = True
|
| 115 |
+
_tied_weights_keys = ["decoder.embed_tokens.weight", "output_projection.weight"]
|
| 116 |
+
_no_split_modules = ["ConformerBlock", "BartDecoderLayer"]
|
| 117 |
+
|
| 118 |
+
def __init__(self, config):
|
| 119 |
+
super().__init__(config)
|
| 120 |
+
self.hparams = SimpleNamespace(input_audio_length=config.input_audio_length, time_hz=config.time_hz)
|
| 121 |
+
self.max_output_seq_len = config.max_output_seq_len
|
| 122 |
+
self.tokenizer = SheetSage2Tokenizer(
|
| 123 |
+
config.input_audio_length, config.time_hz, config.tokenizer_schema_version,
|
| 124 |
+
expected_fingerprint=config.tokenizer_fingerprint,
|
| 125 |
+
)
|
| 126 |
+
if self.tokenizer.n_tokens != config.vocab_size:
|
| 127 |
+
raise ValueError("Tokenizer vocabulary size does not match the model.")
|
| 128 |
+
if config.weights_format == "adapter":
|
| 129 |
+
self.encoder = None
|
| 130 |
+
self.adapter = EncoderAdapters(config)
|
| 131 |
+
else:
|
| 132 |
+
ec = MERT2Config(**config.backbone_config)
|
| 133 |
+
ec._attn_implementation = config.encoder_attn_implementation
|
| 134 |
+
self.encoder = MERT2Model(ec)
|
| 135 |
+
self.adapter = None
|
| 136 |
+
self.layer_weight = nn.Parameter(torch.zeros(config.backbone_config["num_hidden_layers"] + 1))
|
| 137 |
+
self.encoder_projection = nn.Linear(config.backbone_config["hidden_size"], config.hidden_size)
|
| 138 |
+
dc = BartConfig(
|
| 139 |
+
vocab_size=config.vocab_size, d_model=config.hidden_size,
|
| 140 |
+
decoder_layers=config.decoder_layers, decoder_attention_heads=config.num_attention_heads,
|
| 141 |
+
decoder_ffn_dim=config.intermediate_size, max_position_embeddings=config.max_output_seq_len,
|
| 142 |
+
dropout=config.decoder_dropout, attention_dropout=config.decoder_dropout,
|
| 143 |
+
activation_dropout=config.decoder_dropout, activation_function="gelu",
|
| 144 |
+
pad_token_id=config.pad_token_id, bos_token_id=config.bos_token_id,
|
| 145 |
+
eos_token_id=config.eos_token_id, is_encoder_decoder=True, use_cache=True,
|
| 146 |
+
)
|
| 147 |
+
dc._attn_implementation = "sdpa"
|
| 148 |
+
self.token_embedding = nn.Embedding(config.vocab_size, config.hidden_size, padding_idx=config.pad_token_id)
|
| 149 |
+
self.decoder = BartDecoder(dc, embed_tokens=self.token_embedding)
|
| 150 |
+
self.decoder.gradient_checkpointing_disable()
|
| 151 |
+
self.output_projection = nn.Linear(config.hidden_size, config.vocab_size, bias=False)
|
| 152 |
+
self.output_projection.weight = self.token_embedding.weight
|
| 153 |
+
self.lora_merged = config.weights_format == "merged"
|
| 154 |
+
self.post_init()
|
| 155 |
+
|
| 156 |
+
def get_input_embeddings(self):
|
| 157 |
+
return self.token_embedding
|
| 158 |
+
|
| 159 |
+
def set_input_embeddings(self, value):
|
| 160 |
+
self.token_embedding = value
|
| 161 |
+
self.decoder.embed_tokens.weight = value.weight
|
| 162 |
+
|
| 163 |
+
def tie_weights(self):
|
| 164 |
+
super().tie_weights()
|
| 165 |
+
if hasattr(self, "decoder") and hasattr(self, "token_embedding"):
|
| 166 |
+
self.decoder.embed_tokens.weight = self.token_embedding.weight
|
| 167 |
+
|
| 168 |
+
def get_output_embeddings(self):
|
| 169 |
+
return self.output_projection
|
| 170 |
+
|
| 171 |
+
def set_output_embeddings(self, value):
|
| 172 |
+
self.output_projection = value
|
| 173 |
+
|
| 174 |
+
def _init_weights(self, module):
|
| 175 |
+
if isinstance(module, (nn.Linear, nn.Embedding)):
|
| 176 |
+
nn.init.normal_(module.weight, mean=0.0, std=0.02)
|
| 177 |
+
if getattr(module, "bias", None) is not None:
|
| 178 |
+
nn.init.zeros_(module.bias)
|
| 179 |
+
if isinstance(module, nn.Embedding) and module.padding_idx is not None:
|
| 180 |
+
module.weight.data[module.padding_idx].zero_()
|
| 181 |
+
elif isinstance(module, nn.LayerNorm):
|
| 182 |
+
nn.init.ones_(module.weight)
|
| 183 |
+
nn.init.zeros_(module.bias)
|
| 184 |
+
|
| 185 |
+
@torch.no_grad()
|
| 186 |
+
def merge_lora(self):
|
| 187 |
+
if self.lora_merged:
|
| 188 |
+
return self
|
| 189 |
+
if self.encoder is None or self.adapter is None:
|
| 190 |
+
raise RuntimeError("Load both the MERT-v2 parent and adapters before merging.")
|
| 191 |
+
scale = self.config.lora_alpha / self.config.lora_rank
|
| 192 |
+
with torch.autocast("cpu", enabled=False):
|
| 193 |
+
for layer, adapter in zip(self.encoder.layers, self.adapter.layers):
|
| 194 |
+
for name in PROJECTIONS:
|
| 195 |
+
projection = getattr(layer.attn, name)
|
| 196 |
+
update = getattr(adapter.attn, name)
|
| 197 |
+
values = (projection.weight, update.lora_A.weight, update.lora_B.weight)
|
| 198 |
+
if any(value.device.type != "cpu" or value.dtype != torch.float32 for value in values):
|
| 199 |
+
raise ValueError("Merge adapters on CPU with float32 parameters.")
|
| 200 |
+
projection.weight.add_((update.lora_B.weight @ update.lora_A.weight) * scale)
|
| 201 |
+
self.adapter = None
|
| 202 |
+
self.lora_merged = True
|
| 203 |
+
self.config.weights_format = "merged"
|
| 204 |
+
return self
|
| 205 |
+
|
| 206 |
+
@classmethod
|
| 207 |
+
def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs):
|
| 208 |
+
base_path = kwargs.pop("base_model_path", None)
|
| 209 |
+
requested_dtype = kwargs.pop("torch_dtype", torch.float32)
|
| 210 |
+
target_device = kwargs.pop("device_map", None)
|
| 211 |
+
backend = kwargs.pop("attn_implementation", None)
|
| 212 |
+
return_loading_info = kwargs.pop("output_loading_info", False)
|
| 213 |
+
config = kwargs.pop("config", None)
|
| 214 |
+
hub_keys = ("cache_dir", "force_download", "local_files_only", "token", "revision", "subfolder")
|
| 215 |
+
hub_args = {name: kwargs[name] for name in hub_keys if name in kwargs}
|
| 216 |
+
if config is None:
|
| 217 |
+
config = cls.config_class.from_pretrained(pretrained_model_name_or_path, **hub_args)
|
| 218 |
+
if backend is not None:
|
| 219 |
+
config.encoder_attn_implementation = backend
|
| 220 |
+
if config.encoder_attn_implementation not in {"sdpa", "flash_attention_2"}:
|
| 221 |
+
raise ValueError("Select attn_implementation='sdpa' or 'flash_attention_2'.")
|
| 222 |
+
if isinstance(target_device, dict):
|
| 223 |
+
if set(target_device) != {""}:
|
| 224 |
+
raise ValueError("Use a single device for SheetSage2: device_map={'': 'cuda:0'}.")
|
| 225 |
+
target_device = target_device[""]
|
| 226 |
+
if target_device == "auto":
|
| 227 |
+
target_device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 228 |
+
if requested_dtype == "auto":
|
| 229 |
+
requested_dtype = torch.float32
|
| 230 |
+
if isinstance(requested_dtype, str):
|
| 231 |
+
requested_dtype = getattr(torch, requested_dtype, None)
|
| 232 |
+
if requested_dtype not in {None, torch.float32, torch.bfloat16, torch.float16}:
|
| 233 |
+
raise ValueError("Use torch_dtype float32, bfloat16, float16, or 'auto'.")
|
| 234 |
+
# Adapters must be loaded and merged before any reduced-precision cast.
|
| 235 |
+
model, loading_info = super().from_pretrained(
|
| 236 |
+
pretrained_model_name_or_path, *model_args, config=config,
|
| 237 |
+
torch_dtype=torch.float32, attn_implementation="sdpa", output_loading_info=True, **kwargs,
|
| 238 |
+
)
|
| 239 |
+
if any(loading_info.get(name) for name in ("missing_keys", "unexpected_keys", "mismatched_keys", "error_msgs")):
|
| 240 |
+
raise ValueError(f"Incomplete or incompatible SheetSage2 weights: {loading_info}")
|
| 241 |
+
if config.weights_format == "adapter":
|
| 242 |
+
parent = str(base_path or config.base_model_name_or_path)
|
| 243 |
+
parent_hub_args = {name: hub_args[name] for name in ("cache_dir", "force_download", "local_files_only", "token") if name in hub_args}
|
| 244 |
+
parent_hub_args["revision"] = config.base_model_revision
|
| 245 |
+
expected_files = dict(BASE_CODE_HASHES, **{"model.safetensors": config.base_model_sha256})
|
| 246 |
+
for filename, expected in expected_files.items():
|
| 247 |
+
resolved = cached_file(parent, filename, **parent_hub_args)
|
| 248 |
+
if _sha256(resolved) != expected:
|
| 249 |
+
raise ValueError(f"MERT-v2 parent integrity check failed: {filename}")
|
| 250 |
+
model.encoder = AutoModel.from_pretrained(
|
| 251 |
+
parent, trust_remote_code=True, code_revision=config.base_model_revision,
|
| 252 |
+
torch_dtype=torch.float32, attn_implementation=config.encoder_attn_implementation,
|
| 253 |
+
**parent_hub_args,
|
| 254 |
+
)
|
| 255 |
+
for name, expected in config.backbone_config.items():
|
| 256 |
+
if name in ("hidden_size", "intermediate_size", "num_hidden_layers", "num_attention_heads",
|
| 257 |
+
"sampling_rate", "hop_length", "n_fft", "win_length", "num_mel_bins",
|
| 258 |
+
"conv_depthwise_kernel_size", "rotary_embedding_base", "subsampling_channels",
|
| 259 |
+
"subsampling_depths", "layer_norm_eps", "subsampling_layer_norm_eps", "variant"):
|
| 260 |
+
if getattr(model.encoder.config, name) != expected:
|
| 261 |
+
raise ValueError(f"MERT-v2 parent architecture mismatch: {name}")
|
| 262 |
+
model.merge_lora()
|
| 263 |
+
if requested_dtype is not None:
|
| 264 |
+
model.to(dtype=requested_dtype)
|
| 265 |
+
if target_device is not None:
|
| 266 |
+
model.to(target_device)
|
| 267 |
+
# Resource lookup follows the caller's cache and offline settings. These
|
| 268 |
+
# transient options are never written into config or model weights.
|
| 269 |
+
model._hub_resource_options = {name: hub_args[name] for name in ("cache_dir", "local_files_only", "token") if name in hub_args}
|
| 270 |
+
resource_args = dict(model._hub_resource_options, local_files_only=True,
|
| 271 |
+
revision=config._commit_hash or hub_args.get("revision"),
|
| 272 |
+
subfolder=hub_args.get("subfolder", ""),
|
| 273 |
+
_raise_exceptions_for_missing_entries=False)
|
| 274 |
+
resource_config = cached_file(str(pretrained_model_name_or_path), "config.json", **resource_args)
|
| 275 |
+
model._source_snapshot = Path(resource_config).parent if resource_config else Path(pretrained_model_name_or_path)
|
| 276 |
+
model.eval().requires_grad_(False)
|
| 277 |
+
return (model, loading_info) if return_loading_info else model
|
| 278 |
+
|
| 279 |
+
def save_pretrained(self, save_directory, *args, **kwargs):
|
| 280 |
+
if not self.lora_merged or self.encoder is None:
|
| 281 |
+
raise ValueError("Load and merge the model before saving a standalone snapshot.")
|
| 282 |
+
original_config = self.config
|
| 283 |
+
source = original_config._name_or_path
|
| 284 |
+
self.config = copy.deepcopy(original_config)
|
| 285 |
+
self.config.weights_format = "merged"
|
| 286 |
+
self.config._name_or_path = ""
|
| 287 |
+
self.config.backbone_config.pop("_name_or_path", None)
|
| 288 |
+
try:
|
| 289 |
+
result = super().save_pretrained(save_directory, *args, **kwargs)
|
| 290 |
+
from .processing_sheetsage2 import SheetSage2Processor
|
| 291 |
+
SheetSage2Processor.from_model_config(self.config).save_pretrained(save_directory)
|
| 292 |
+
finally:
|
| 293 |
+
self.config = original_config
|
| 294 |
+
source_dir = getattr(self, "_source_snapshot", Path(source))
|
| 295 |
+
if source and not source_dir.is_dir():
|
| 296 |
+
from huggingface_hub import snapshot_download
|
| 297 |
+
try:
|
| 298 |
+
source_dir = Path(snapshot_download(source, revision=original_config._commit_hash,
|
| 299 |
+
cache_dir=getattr(self, "_hub_resource_options", {}).get("cache_dir"),
|
| 300 |
+
local_files_only=True))
|
| 301 |
+
except (OSError, ValueError):
|
| 302 |
+
source_dir = None
|
| 303 |
+
if source and source_dir is not None and source_dir.is_dir():
|
| 304 |
+
destination = Path(save_directory)
|
| 305 |
+
for name in ("infer.py", "render.py", "setup_render.py", "requirements.txt", "requirements-render.txt",
|
| 306 |
+
"LICENSE", "THIRD_PARTY_NOTICES.md"):
|
| 307 |
+
if (source_dir / name).is_file() and (source_dir / name).resolve() != (destination / name).resolve():
|
| 308 |
+
shutil.copy2(source_dir / name, destination / name)
|
| 309 |
+
assets = source_dir / "render_assets"
|
| 310 |
+
if (assets / "manifest.json").is_file() and assets.resolve() != (destination / "render_assets").resolve():
|
| 311 |
+
from .rendering_sheetsage2 import _verify_assets
|
| 312 |
+
try:
|
| 313 |
+
_verify_assets(assets, audio=True)
|
| 314 |
+
except (FileNotFoundError, ValueError):
|
| 315 |
+
pass
|
| 316 |
+
else:
|
| 317 |
+
shutil.copytree(assets, destination / "render_assets", dirs_exist_ok=True)
|
| 318 |
+
return result
|
| 319 |
+
|
| 320 |
+
def _prepare_audio(self, input_values, attention_mask=None):
|
| 321 |
+
if self.encoder is None or not self.lora_merged:
|
| 322 |
+
raise RuntimeError("Load the model with from_pretrained before inference.")
|
| 323 |
+
if input_values.ndim != 2 or input_values.shape[0] < 1 or not input_values.is_floating_point():
|
| 324 |
+
raise ValueError("input_values must be floating-point [batch, samples].")
|
| 325 |
+
if not torch.isfinite(input_values).all():
|
| 326 |
+
raise ValueError("Audio must contain only finite samples.")
|
| 327 |
+
batch, samples = input_values.shape
|
| 328 |
+
minimum = self.encoder.config.minimum_input_samples
|
| 329 |
+
window = round(self.config.input_audio_length * self.config.sampling_rate)
|
| 330 |
+
if samples < minimum:
|
| 331 |
+
raise ValueError(f"Each waveform must contain at least {minimum} samples.")
|
| 332 |
+
if samples > window:
|
| 333 |
+
raise ValueError("Audio exceeds one model window; use transcribe for whole songs.")
|
| 334 |
+
if attention_mask is None:
|
| 335 |
+
lengths = torch.full((batch,), samples, dtype=torch.long, device=input_values.device)
|
| 336 |
+
else:
|
| 337 |
+
if attention_mask.shape != input_values.shape:
|
| 338 |
+
raise ValueError("attention_mask must have the same shape as input_values.")
|
| 339 |
+
mask = attention_mask.to(device=input_values.device)
|
| 340 |
+
if not ((mask == 0) | (mask == 1)).all():
|
| 341 |
+
raise ValueError("attention_mask must contain zeros and ones.")
|
| 342 |
+
mask = mask.bool()
|
| 343 |
+
lengths = mask.sum(1)
|
| 344 |
+
expected = torch.arange(samples, device=input_values.device)[None] < lengths[:, None]
|
| 345 |
+
if not torch.equal(mask, expected) or (lengths < minimum).any():
|
| 346 |
+
raise ValueError("attention_mask must identify a nonempty right-padded waveform.")
|
| 347 |
+
input_values = input_values.masked_fill(~mask, 0)
|
| 348 |
+
# The encoder attends to the complete fixed window, including this silence.
|
| 349 |
+
input_values = F.pad(input_values.float(), (0, window - samples))
|
| 350 |
+
stride = self.encoder.config.inputs_to_logits_ratio
|
| 351 |
+
if window % stride:
|
| 352 |
+
input_values = F.pad(input_values, (0, stride - window % stride))
|
| 353 |
+
return input_values, lengths
|
| 354 |
+
|
| 355 |
+
def get_audio_features(self, input_values, attention_mask=None, output_hidden_states=None, return_dict=True):
|
| 356 |
+
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
|
| 357 |
+
waveform, lengths = self._prepare_audio(input_values, attention_mask)
|
| 358 |
+
mel = self.encoder.feature_extractor(waveform)
|
| 359 |
+
weight_dtype = self.encoder.subsampling_module[0].convnext_layers[0].depthwise_block[1].weight.dtype
|
| 360 |
+
hidden = self.encoder.subsampling_module(mel.to(dtype=weight_dtype))
|
| 361 |
+
input_hidden = hidden if output_hidden_states else None
|
| 362 |
+
weights = torch.softmax(self.layer_weight, dim=0)
|
| 363 |
+
mixed = hidden * weights[0]
|
| 364 |
+
positions = self.encoder.embed_positions(hidden)
|
| 365 |
+
states = [] if output_hidden_states else None
|
| 366 |
+
for weight, layer in zip(weights[1:], self.encoder.layers):
|
| 367 |
+
hidden = layer(hidden, positions)
|
| 368 |
+
mixed = mixed + hidden * weight
|
| 369 |
+
if states is not None:
|
| 370 |
+
states.append(hidden)
|
| 371 |
+
memory = self.encoder_projection(mixed)
|
| 372 |
+
stride = self.encoder.config.inputs_to_logits_ratio
|
| 373 |
+
frame_mask = torch.arange(memory.shape[1], device=memory.device)[None] < ((lengths + stride - 1) // stride)[:, None]
|
| 374 |
+
output = SheetSage2EncoderOutput(
|
| 375 |
+
encoder_last_hidden_state=memory, backbone_last_hidden_state=hidden,
|
| 376 |
+
mixed_hidden_state=mixed, input_hidden_state=input_hidden,
|
| 377 |
+
backbone_hidden_states=tuple(states) if states is not None else None,
|
| 378 |
+
feature_attention_mask=frame_mask,
|
| 379 |
+
)
|
| 380 |
+
return output if return_dict else output.to_tuple()
|
| 381 |
+
|
| 382 |
+
def encode(self, audio):
|
| 383 |
+
return self.get_audio_features(audio).last_hidden_state
|
| 384 |
+
|
| 385 |
+
def _decode(self, memory, decoder_input_ids, use_cache=False, past_key_values=None, output_hidden_states=False):
|
| 386 |
+
attention_mask = None if past_key_values is not None else decoder_input_ids != self.tokenizer.pad_token
|
| 387 |
+
output = self.decoder(
|
| 388 |
+
input_ids=decoder_input_ids, attention_mask=attention_mask,
|
| 389 |
+
encoder_hidden_states=memory, encoder_attention_mask=None,
|
| 390 |
+
past_key_values=past_key_values, use_cache=use_cache,
|
| 391 |
+
output_hidden_states=output_hidden_states, return_dict=True,
|
| 392 |
+
)
|
| 393 |
+
return self.output_projection(output.last_hidden_state), output
|
| 394 |
+
|
| 395 |
+
def decode(self, memory, decoder_input_ids, use_cache=False, past_key_values=None):
|
| 396 |
+
logits, output = self._decode(memory, decoder_input_ids, use_cache, past_key_values)
|
| 397 |
+
return logits, output.past_key_values
|
| 398 |
+
|
| 399 |
+
def forward(self, input_values=None, decoder_input_ids=None, attention_mask=None,
|
| 400 |
+
encoder_outputs=None, past_key_values=None, use_cache=None,
|
| 401 |
+
output_hidden_states=None, return_dict=None):
|
| 402 |
+
output_hidden_states = self.config.output_hidden_states if output_hidden_states is None else output_hidden_states
|
| 403 |
+
use_cache = self.config.use_cache if use_cache is None else use_cache
|
| 404 |
+
if decoder_input_ids is None:
|
| 405 |
+
raise ValueError("decoder_input_ids is required for logits; use generate to transcribe audio.")
|
| 406 |
+
if encoder_outputs is None:
|
| 407 |
+
if input_values is None:
|
| 408 |
+
raise ValueError("Provide input_values or encoder_outputs.")
|
| 409 |
+
encoder_outputs = self.get_audio_features(input_values, attention_mask, output_hidden_states)
|
| 410 |
+
if torch.is_tensor(encoder_outputs):
|
| 411 |
+
encoder_outputs = SheetSage2EncoderOutput(encoder_last_hidden_state=encoder_outputs)
|
| 412 |
+
logits, decoded = self._decode(encoder_outputs.last_hidden_state, decoder_input_ids,
|
| 413 |
+
use_cache, past_key_values, output_hidden_states)
|
| 414 |
+
output = SheetSage2Output(
|
| 415 |
+
logits=logits, past_key_values=decoded.past_key_values,
|
| 416 |
+
decoder_hidden_states=decoded.hidden_states,
|
| 417 |
+
encoder_last_hidden_state=encoder_outputs.last_hidden_state,
|
| 418 |
+
backbone_last_hidden_state=encoder_outputs.backbone_last_hidden_state if output_hidden_states else None,
|
| 419 |
+
mixed_hidden_state=encoder_outputs.mixed_hidden_state if output_hidden_states else None,
|
| 420 |
+
input_hidden_state=encoder_outputs.input_hidden_state if output_hidden_states else None,
|
| 421 |
+
backbone_hidden_states=encoder_outputs.backbone_hidden_states if output_hidden_states else None,
|
| 422 |
+
feature_attention_mask=encoder_outputs.feature_attention_mask,
|
| 423 |
+
)
|
| 424 |
+
return_dict = self.config.use_return_dict if return_dict is None else return_dict
|
| 425 |
+
return output if return_dict else output.to_tuple()
|
| 426 |
+
|
| 427 |
+
def generate(self, input_values, **kwargs):
|
| 428 |
+
"""Generate grammar-constrained symbolic tokens with autoregressive caching."""
|
| 429 |
+
from .generation_sheetsage2 import generate
|
| 430 |
+
return generate(self, input_values, **kwargs)
|
| 431 |
+
|
| 432 |
+
def transcribe(self, audio, output_dir=None, *, melody_only=False, **kwargs):
|
| 433 |
+
"""Return transcription in memory; set output_dir to also save files.
|
| 434 |
+
|
| 435 |
+
Accepts a path, encoded audio bytes, binary stream, or waveform with
|
| 436 |
+
sampling_rate. Returns ABC text, MIDI bytes, timed events, and optional
|
| 437 |
+
per-window CPU tensors. With output_dir, optional tensors are saved
|
| 438 |
+
instead of retained in memory. Set melody_only=True to retain both vocal
|
| 439 |
+
and instrumental melodies while omitting chords from ABC and playback;
|
| 440 |
+
raw predicted annotations remain available. The default keeps full
|
| 441 |
+
transcription. See the model card for rendering options.
|
| 442 |
+
"""
|
| 443 |
+
from .pipeline_sheetsage2 import transcribe
|
| 444 |
+
return transcribe(self, audio, output_dir=output_dir, melody_only=melody_only, **kwargs)
|
| 445 |
+
|
| 446 |
+
|
| 447 |
+
SheetSage2ForConditionalGeneration = SheetSage2Model
|
| 448 |
+
SheetSage2Model.register_for_auto_class("AutoModel")
|
notation_sheetsage2.py
ADDED
|
@@ -0,0 +1,1570 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Convert timed musical events into validated two-voice ABC notation."""
|
| 2 |
+
|
| 3 |
+
from __future__ import annotations
|
| 4 |
+
|
| 5 |
+
import math
|
| 6 |
+
import os
|
| 7 |
+
import re
|
| 8 |
+
from collections import Counter
|
| 9 |
+
from dataclasses import dataclass, replace
|
| 10 |
+
from io import BytesIO
|
| 11 |
+
from pathlib import Path
|
| 12 |
+
from typing import Iterable, Sequence
|
| 13 |
+
|
| 14 |
+
import numpy as np
|
| 15 |
+
import pretty_midi
|
| 16 |
+
|
| 17 |
+
from .io_sheetsage2 import atomic_write_text
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
SUBBEAT_DIVISION = 4
|
| 21 |
+
VOICE_IDS = ("Vocal", "Ins")
|
| 22 |
+
NO_CHORDS = frozenset({"N", "X", "?"})
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
class AbcRebuildError(ValueError):
|
| 26 |
+
"""Base class for deterministic reconstruction failures."""
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
class BeatGridError(AbcRebuildError):
|
| 30 |
+
pass
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
class ChordSymbolError(AbcRebuildError):
|
| 34 |
+
pass
|
| 35 |
+
|
| 36 |
+
|
| 37 |
+
class MelodyVoiceError(AbcRebuildError):
|
| 38 |
+
pass
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
@dataclass(frozen=True)
|
| 42 |
+
class BeatEvent:
|
| 43 |
+
time: float
|
| 44 |
+
beat_id: int
|
| 45 |
+
declared_numerator: int
|
| 46 |
+
denominator: int
|
| 47 |
+
line_no: int
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
@dataclass(frozen=True)
|
| 51 |
+
class Measure:
|
| 52 |
+
index: int
|
| 53 |
+
start_beat: int
|
| 54 |
+
end_beat: int
|
| 55 |
+
numerator: int
|
| 56 |
+
denominator: int
|
| 57 |
+
pickup: bool = False
|
| 58 |
+
partial: bool = False
|
| 59 |
+
inferred: bool = False
|
| 60 |
+
notated_numerator: int | None = None
|
| 61 |
+
notated_denominator: int | None = None
|
| 62 |
+
pad_before: bool = False
|
| 63 |
+
|
| 64 |
+
@property
|
| 65 |
+
def beat_count(self) -> int:
|
| 66 |
+
return self.end_beat - self.start_beat
|
| 67 |
+
|
| 68 |
+
@property
|
| 69 |
+
def start_t(self) -> int:
|
| 70 |
+
return self.start_beat * SUBBEAT_DIVISION
|
| 71 |
+
|
| 72 |
+
@property
|
| 73 |
+
def end_t(self) -> int:
|
| 74 |
+
return self.end_beat * SUBBEAT_DIVISION
|
| 75 |
+
|
| 76 |
+
@property
|
| 77 |
+
def abc_numerator(self) -> int:
|
| 78 |
+
return self.notated_numerator or self.numerator
|
| 79 |
+
|
| 80 |
+
@property
|
| 81 |
+
def abc_denominator(self) -> int:
|
| 82 |
+
return self.notated_denominator or self.denominator
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
@dataclass
|
| 86 |
+
class RebuiltAbcScore:
|
| 87 |
+
beats: list[BeatEvent]
|
| 88 |
+
measures: list[Measure]
|
| 89 |
+
subbeat_times: np.ndarray
|
| 90 |
+
subbeat_quarters: np.ndarray
|
| 91 |
+
subbeat_denominators: np.ndarray
|
| 92 |
+
key_arr: np.ndarray
|
| 93 |
+
chord_arr: np.ndarray
|
| 94 |
+
structure_events: list[tuple[int, str]]
|
| 95 |
+
voice_arrs: dict[str, np.ndarray]
|
| 96 |
+
diagnostics: list[str]
|
| 97 |
+
subbeat_div: int = SUBBEAT_DIVISION
|
| 98 |
+
|
| 99 |
+
|
| 100 |
+
@dataclass
|
| 101 |
+
class MeasureGroup:
|
| 102 |
+
measures: list[Measure]
|
| 103 |
+
structure_labels: list[str]
|
| 104 |
+
meter_changed: bool
|
| 105 |
+
key_changed: bool
|
| 106 |
+
|
| 107 |
+
|
| 108 |
+
_QUALITY_TO_ABC = {
|
| 109 |
+
"maj": "",
|
| 110 |
+
"min": "m",
|
| 111 |
+
"dim": "dim",
|
| 112 |
+
"aug": "aug",
|
| 113 |
+
"7": "7",
|
| 114 |
+
"maj7": "maj7",
|
| 115 |
+
"min7": "m7",
|
| 116 |
+
"dim7": "dim7",
|
| 117 |
+
"hdim7": "m7b5",
|
| 118 |
+
"sus4": "sus4",
|
| 119 |
+
"sus2": "sus2",
|
| 120 |
+
"maj6": "6",
|
| 121 |
+
"min6": "m6",
|
| 122 |
+
"sus4(b7)": "7sus4",
|
| 123 |
+
# abc2midi and SymMusic both accept the parenthesized major seventh.
|
| 124 |
+
# Common aliases such as mmaj7/mM7 trigger abc2midi diagnostics.
|
| 125 |
+
"minmaj7": "m(maj7)",
|
| 126 |
+
}
|
| 127 |
+
|
| 128 |
+
_NATURAL_PITCH_CLASS = {
|
| 129 |
+
"C": 0,
|
| 130 |
+
"D": 2,
|
| 131 |
+
"E": 4,
|
| 132 |
+
"F": 5,
|
| 133 |
+
"G": 7,
|
| 134 |
+
"A": 9,
|
| 135 |
+
"B": 11,
|
| 136 |
+
}
|
| 137 |
+
_LETTERS = "CDEFGAB"
|
| 138 |
+
_SHARP_PITCH_NAMES = ("C", "C#", "D", "D#", "E", "F", "F#", "G", "G#", "A", "A#", "B")
|
| 139 |
+
_FLAT_PITCH_NAMES = ("C", "Db", "D", "Eb", "E", "F", "Gb", "G", "Ab", "A", "Bb", "B")
|
| 140 |
+
_ROOT_RE = re.compile(r"^(?P<letter>[A-G])(?P<accidental>#{0,2}|b{0,2})$")
|
| 141 |
+
_BASS_DEGREE_RE = re.compile(r"^(?P<accidental>#{0,2}|b{0,2})(?P<degree>[1-9]|1[0-3])$")
|
| 142 |
+
|
| 143 |
+
_KEY_SIGNATURE_ACCIDENTALS = {
|
| 144 |
+
"C": 0,
|
| 145 |
+
"G": 1,
|
| 146 |
+
"D": 2,
|
| 147 |
+
"A": 3,
|
| 148 |
+
"E": 4,
|
| 149 |
+
"B": 5,
|
| 150 |
+
"F#": 6,
|
| 151 |
+
"C#": 7,
|
| 152 |
+
"F": -1,
|
| 153 |
+
"Bb": -2,
|
| 154 |
+
"Eb": -3,
|
| 155 |
+
"Ab": -4,
|
| 156 |
+
"Db": -5,
|
| 157 |
+
"Gb": -6,
|
| 158 |
+
"Cb": -7,
|
| 159 |
+
"Am": 0,
|
| 160 |
+
"Em": 1,
|
| 161 |
+
"Bm": 2,
|
| 162 |
+
"F#m": 3,
|
| 163 |
+
"C#m": 4,
|
| 164 |
+
"G#m": 5,
|
| 165 |
+
"D#m": 6,
|
| 166 |
+
"A#m": 7,
|
| 167 |
+
"Dm": -1,
|
| 168 |
+
"Gm": -2,
|
| 169 |
+
"Cm": -3,
|
| 170 |
+
"Fm": -4,
|
| 171 |
+
"Bbm": -5,
|
| 172 |
+
"Ebm": -6,
|
| 173 |
+
"Abm": -7,
|
| 174 |
+
}
|
| 175 |
+
|
| 176 |
+
# Keep the standard key-relative chromatic spelling. MIDI carries only
|
| 177 |
+
# pitch, not note names, so a deterministic key-based table is preferable to
|
| 178 |
+
# rewriting every chromatic pitch with a single sharp/flat. In remote sharp
|
| 179 |
+
# and flat keys this deliberately permits musically useful double accidentals,
|
| 180 |
+
# for example MIDI G as F## in G# minor.
|
| 181 |
+
_KEY_RELATIVE_PITCH_NAMES = {
|
| 182 |
+
7: ("B#", "C#", "C##", "D#", "D##", "E#", "F#", "F##", "G#", "G##", "A#", "B"),
|
| 183 |
+
6: ("B#", "C#", "C##", "D#", "E", "E#", "F#", "F##", "G#", "G##", "A#", "B"),
|
| 184 |
+
5: ("B#", "C#", "C##", "D#", "E", "E#", "F#", "F##", "G#", "A", "A#", "B"),
|
| 185 |
+
4: ("B#", "C#", "D", "D#", "E", "E#", "F#", "F##", "G#", "A", "A#", "B"),
|
| 186 |
+
3: ("B#", "C#", "D", "D#", "E", "E#", "F#", "G", "G#", "A", "A#", "B"),
|
| 187 |
+
2: ("C", "C#", "D", "D#", "E", "E#", "F#", "G", "G#", "A", "A#", "B"),
|
| 188 |
+
1: ("C", "C#", "D", "D#", "E", "F", "F#", "G", "G#", "A", "A#", "B"),
|
| 189 |
+
0: ("C", "C#", "D", "D#", "E", "F", "F#", "G", "G#", "A", "Bb", "B"),
|
| 190 |
+
-1: ("C", "C#", "D", "Eb", "E", "F", "F#", "G", "G#", "A", "Bb", "B"),
|
| 191 |
+
-2: ("C", "C#", "D", "Eb", "E", "F", "F#", "G", "Ab", "A", "Bb", "B"),
|
| 192 |
+
-3: ("C", "Db", "D", "Eb", "E", "F", "F#", "G", "Ab", "A", "Bb", "B"),
|
| 193 |
+
-4: ("C", "Db", "D", "Eb", "E", "F", "Gb", "G", "Ab", "A", "Bb", "B"),
|
| 194 |
+
-5: ("C", "Db", "D", "Eb", "E", "F", "Gb", "G", "Ab", "A", "Bb", "Cb"),
|
| 195 |
+
-6: ("C", "Db", "D", "Eb", "Fb", "F", "Gb", "G", "Ab", "A", "Bb", "Cb"),
|
| 196 |
+
-7: ("C", "Db", "D", "Eb", "Fb", "F", "Gb", "G", "Ab", "Bbb", "Bb", "Cb"),
|
| 197 |
+
}
|
| 198 |
+
|
| 199 |
+
|
| 200 |
+
def _read_tsv(path: os.PathLike[str] | str, min_columns: int) -> list[tuple[int, list[str]]]:
|
| 201 |
+
rows = []
|
| 202 |
+
with open(path, "r", encoding="utf-8-sig") as handle:
|
| 203 |
+
for line_no, raw_line in enumerate(handle, 1):
|
| 204 |
+
line = raw_line.rstrip("\r\n")
|
| 205 |
+
if not line.strip():
|
| 206 |
+
continue
|
| 207 |
+
columns = line.split("\t")
|
| 208 |
+
if len(columns) < min_columns:
|
| 209 |
+
raise AbcRebuildError(
|
| 210 |
+
f"{path}:{line_no}: expected at least {min_columns} tab-separated columns"
|
| 211 |
+
)
|
| 212 |
+
rows.append((line_no, columns))
|
| 213 |
+
return rows
|
| 214 |
+
|
| 215 |
+
|
| 216 |
+
def read_beats(path: os.PathLike[str] | str) -> list[BeatEvent]:
|
| 217 |
+
return _parse_beats(_read_tsv(path, 3), path)
|
| 218 |
+
|
| 219 |
+
|
| 220 |
+
def _row_entries(rows, source):
|
| 221 |
+
entries = []
|
| 222 |
+
for line_no, row in enumerate(rows, 1):
|
| 223 |
+
if len(row) < 3:
|
| 224 |
+
raise AbcRebuildError(f"{source}:{line_no}: expected at least 3 columns")
|
| 225 |
+
entries.append((line_no, [str(value) for value in row]))
|
| 226 |
+
return entries
|
| 227 |
+
|
| 228 |
+
|
| 229 |
+
def _parse_beats(entries, path):
|
| 230 |
+
beats = []
|
| 231 |
+
for line_no, row in entries:
|
| 232 |
+
meter_text = row[2]
|
| 233 |
+
if len(row) >= 4:
|
| 234 |
+
numerator_text, denominator_text = meter_text, row[3]
|
| 235 |
+
elif "/" in meter_text:
|
| 236 |
+
numerator_text, denominator_text = meter_text.split("/", 1)
|
| 237 |
+
else:
|
| 238 |
+
numerator_text, denominator_text = meter_text, "4"
|
| 239 |
+
try:
|
| 240 |
+
beat = BeatEvent(
|
| 241 |
+
time=float(row[0]),
|
| 242 |
+
beat_id=int(row[1]),
|
| 243 |
+
declared_numerator=int(numerator_text),
|
| 244 |
+
denominator=int(denominator_text),
|
| 245 |
+
line_no=line_no,
|
| 246 |
+
)
|
| 247 |
+
except ValueError as exc:
|
| 248 |
+
raise BeatGridError(f"{path}:{line_no}: invalid beat row {row!r}") from exc
|
| 249 |
+
if beat.beat_id < 1:
|
| 250 |
+
raise BeatGridError(f"{path}:{line_no}: beat ID must be positive")
|
| 251 |
+
if beat.declared_numerator < 1:
|
| 252 |
+
raise BeatGridError(f"{path}:{line_no}: meter numerator must be positive")
|
| 253 |
+
if beat.denominator < 1 or beat.denominator & (beat.denominator - 1):
|
| 254 |
+
raise BeatGridError(
|
| 255 |
+
f"{path}:{line_no}: meter denominator must be a positive power of two"
|
| 256 |
+
)
|
| 257 |
+
if beats and beat.time <= beats[-1].time:
|
| 258 |
+
raise BeatGridError(f"{path}:{line_no}: beat times must be strictly increasing")
|
| 259 |
+
beats.append(beat)
|
| 260 |
+
if len(beats) < 2:
|
| 261 |
+
raise BeatGridError(f"{path}: at least two beat events are required")
|
| 262 |
+
return beats
|
| 263 |
+
|
| 264 |
+
|
| 265 |
+
def read_chords(path: os.PathLike[str] | str) -> list[tuple[float, float, str]]:
|
| 266 |
+
return _parse_chords(_read_tsv(path, 3), path)
|
| 267 |
+
|
| 268 |
+
|
| 269 |
+
def _parse_chords(entries, path):
|
| 270 |
+
rows = []
|
| 271 |
+
previous_end = None
|
| 272 |
+
for line_no, row in entries:
|
| 273 |
+
start, end, chord = float(row[0]), float(row[1]), row[2].strip()
|
| 274 |
+
if end <= start:
|
| 275 |
+
raise ChordSymbolError(f"{path}:{line_no}: chord end must be after start")
|
| 276 |
+
if previous_end is not None and start < previous_end - 1e-6:
|
| 277 |
+
raise ChordSymbolError(f"{path}:{line_no}: overlapping chord intervals")
|
| 278 |
+
chord_symbol_to_abc(chord)
|
| 279 |
+
rows.append((start, end, chord))
|
| 280 |
+
previous_end = end
|
| 281 |
+
return rows
|
| 282 |
+
|
| 283 |
+
|
| 284 |
+
def read_keys(path: os.PathLike[str] | str) -> list[tuple[float, float, str]]:
|
| 285 |
+
return _parse_keys(_read_tsv(path, 3), path)
|
| 286 |
+
|
| 287 |
+
|
| 288 |
+
def _parse_keys(entries, path):
|
| 289 |
+
rows = []
|
| 290 |
+
previous_end = None
|
| 291 |
+
for line_no, row in entries:
|
| 292 |
+
start, end, key = float(row[0]), float(row[1]), row[2].strip()
|
| 293 |
+
if end <= start:
|
| 294 |
+
raise AbcRebuildError(f"{path}:{line_no}: key end must be after start")
|
| 295 |
+
if previous_end is not None and start < previous_end - 1e-6:
|
| 296 |
+
raise AbcRebuildError(f"{path}:{line_no}: overlapping key intervals")
|
| 297 |
+
normalized = key_symbol_to_abc(key)
|
| 298 |
+
rows.append((start, end, normalized))
|
| 299 |
+
previous_end = end
|
| 300 |
+
if not rows:
|
| 301 |
+
raise AbcRebuildError(f"{path}: at least one key interval is required")
|
| 302 |
+
return rows
|
| 303 |
+
|
| 304 |
+
|
| 305 |
+
def read_structures(path: os.PathLike[str] | str) -> list[tuple[float, float, str]]:
|
| 306 |
+
return _parse_structures(_read_tsv(path, 3), path)
|
| 307 |
+
|
| 308 |
+
|
| 309 |
+
def _parse_structures(entries, path):
|
| 310 |
+
rows = []
|
| 311 |
+
previous_end = None
|
| 312 |
+
for line_no, row in entries:
|
| 313 |
+
start, end, label = float(row[0]), float(row[1]), row[2].strip()
|
| 314 |
+
if end <= start:
|
| 315 |
+
raise AbcRebuildError(f"{path}:{line_no}: structure end must be after start")
|
| 316 |
+
if previous_end is not None and start < previous_end - 1e-6:
|
| 317 |
+
raise AbcRebuildError(f"{path}:{line_no}: overlapping structure intervals")
|
| 318 |
+
rows.append((start, end, label))
|
| 319 |
+
previous_end = end
|
| 320 |
+
return rows
|
| 321 |
+
|
| 322 |
+
|
| 323 |
+
def _mode_with_first_tiebreak(values: Sequence[int]) -> int:
|
| 324 |
+
counts = Counter(values)
|
| 325 |
+
maximum = max(counts.values())
|
| 326 |
+
return next(value for value in values if counts[value] == maximum)
|
| 327 |
+
|
| 328 |
+
|
| 329 |
+
def infer_measures(
|
| 330 |
+
beats: Sequence[BeatEvent],
|
| 331 |
+
*,
|
| 332 |
+
meter_conflict: str = "infer",
|
| 333 |
+
) -> tuple[list[Measure], list[str]]:
|
| 334 |
+
"""Infer self-consistent measures from actual downbeat boundaries."""
|
| 335 |
+
|
| 336 |
+
if meter_conflict not in {"infer", "reject"}:
|
| 337 |
+
raise ValueError("meter_conflict must be 'infer' or 'reject'")
|
| 338 |
+
downbeat_indices = [index for index, beat in enumerate(beats) if beat.beat_id == 1]
|
| 339 |
+
if not downbeat_indices:
|
| 340 |
+
raise BeatGridError("No downbeat (beat ID 1) exists in the beat lab")
|
| 341 |
+
spans: list[tuple[int, int, bool, bool]] = []
|
| 342 |
+
if downbeat_indices[0] > 0:
|
| 343 |
+
spans.append((0, downbeat_indices[0], True, False))
|
| 344 |
+
spans.extend(
|
| 345 |
+
(start, end, False, False)
|
| 346 |
+
for start, end in zip(downbeat_indices, downbeat_indices[1:])
|
| 347 |
+
)
|
| 348 |
+
if downbeat_indices[-1] < len(beats) - 1:
|
| 349 |
+
# Exported beat labs use their last row as the score end boundary. If
|
| 350 |
+
# that row is not a downbeat, the final bar is intentionally truncated.
|
| 351 |
+
spans.append((downbeat_indices[-1], len(beats) - 1, False, True))
|
| 352 |
+
if not spans:
|
| 353 |
+
raise BeatGridError("No positive-length measure exists between downbeats")
|
| 354 |
+
|
| 355 |
+
diagnostics = []
|
| 356 |
+
measures = []
|
| 357 |
+
for measure_index, (start, end, pickup, partial) in enumerate(spans):
|
| 358 |
+
events = list(beats[start:end])
|
| 359 |
+
beat_count = len(events)
|
| 360 |
+
if beat_count < 1:
|
| 361 |
+
raise BeatGridError(f"Measure {measure_index}: empty downbeat span")
|
| 362 |
+
ids = [event.beat_id for event in events]
|
| 363 |
+
expected_ids = list(range(ids[0], ids[0] + beat_count))
|
| 364 |
+
if ids != expected_ids:
|
| 365 |
+
line_numbers = [event.line_no for event in events]
|
| 366 |
+
raise BeatGridError(
|
| 367 |
+
f"Measure {measure_index} (beat rows {line_numbers[0]}-{line_numbers[-1]}): "
|
| 368 |
+
f"non-consecutive beat IDs {ids!r}"
|
| 369 |
+
)
|
| 370 |
+
if not pickup and ids[0] != 1:
|
| 371 |
+
raise BeatGridError(f"Measure {measure_index}: full measure does not start at beat ID 1")
|
| 372 |
+
|
| 373 |
+
denominators = [event.denominator for event in events]
|
| 374 |
+
denominator = _mode_with_first_tiebreak(denominators)
|
| 375 |
+
declared_numerators = [event.declared_numerator for event in events]
|
| 376 |
+
declared_numerator = _mode_with_first_tiebreak(declared_numerators)
|
| 377 |
+
numerator_conflict = any(value != beat_count for value in declared_numerators)
|
| 378 |
+
denominator_conflict = any(value != denominator for value in denominators)
|
| 379 |
+
pad_final_partial = (
|
| 380 |
+
partial
|
| 381 |
+
and len(set(declared_numerators)) == 1
|
| 382 |
+
and not denominator_conflict
|
| 383 |
+
and declared_numerator >= beat_count
|
| 384 |
+
)
|
| 385 |
+
inferred = pickup or partial or numerator_conflict or denominator_conflict
|
| 386 |
+
unresolved_numerator_conflict = (
|
| 387 |
+
numerator_conflict
|
| 388 |
+
and not pad_final_partial
|
| 389 |
+
and not pickup
|
| 390 |
+
)
|
| 391 |
+
if (
|
| 392 |
+
unresolved_numerator_conflict or denominator_conflict
|
| 393 |
+
) and meter_conflict == "reject":
|
| 394 |
+
raise BeatGridError(
|
| 395 |
+
f"Measure {measure_index}: {beat_count} actual beats conflict with declarations "
|
| 396 |
+
f"{list(zip(declared_numerators, denominators))!r}"
|
| 397 |
+
)
|
| 398 |
+
if pad_final_partial and declared_numerator > beat_count:
|
| 399 |
+
diagnostics.append(
|
| 400 |
+
f"measure {measure_index}: padded final {beat_count}/{denominator} span "
|
| 401 |
+
f"to declared {declared_numerator}/{denominator} with trailing rest"
|
| 402 |
+
)
|
| 403 |
+
elif numerator_conflict:
|
| 404 |
+
diagnostics.append(
|
| 405 |
+
f"measure {measure_index}: inferred {beat_count}/{denominator} from downbeat span; "
|
| 406 |
+
f"declared numerators were {declared_numerators}"
|
| 407 |
+
)
|
| 408 |
+
if denominator_conflict:
|
| 409 |
+
diagnostics.append(
|
| 410 |
+
f"measure {measure_index}: placed denominator {denominator} at the measure boundary; "
|
| 411 |
+
f"row declarations were {denominators}"
|
| 412 |
+
)
|
| 413 |
+
measures.append(
|
| 414 |
+
Measure(
|
| 415 |
+
index=measure_index,
|
| 416 |
+
start_beat=start,
|
| 417 |
+
end_beat=end,
|
| 418 |
+
numerator=beat_count,
|
| 419 |
+
denominator=denominator,
|
| 420 |
+
pickup=pickup,
|
| 421 |
+
partial=partial,
|
| 422 |
+
inferred=inferred,
|
| 423 |
+
notated_numerator=(
|
| 424 |
+
declared_numerator if pad_final_partial else beat_count
|
| 425 |
+
),
|
| 426 |
+
)
|
| 427 |
+
)
|
| 428 |
+
if len(measures) >= 2:
|
| 429 |
+
first = measures[0]
|
| 430 |
+
following = measures[1]
|
| 431 |
+
first_duration = first.numerator / first.denominator
|
| 432 |
+
following_duration = (
|
| 433 |
+
following.abc_numerator / following.abc_denominator
|
| 434 |
+
)
|
| 435 |
+
if first_duration < following_duration:
|
| 436 |
+
measures[0] = replace(
|
| 437 |
+
first,
|
| 438 |
+
inferred=True,
|
| 439 |
+
notated_numerator=following.abc_numerator,
|
| 440 |
+
notated_denominator=following.abc_denominator,
|
| 441 |
+
pad_before=True,
|
| 442 |
+
)
|
| 443 |
+
diagnostics.append(
|
| 444 |
+
f"measure 0: padded leading {first.numerator}/{first.denominator} span "
|
| 445 |
+
f"to {following.abc_numerator}/{following.abc_denominator} "
|
| 446 |
+
f"with preceding rest"
|
| 447 |
+
)
|
| 448 |
+
return measures, diagnostics
|
| 449 |
+
|
| 450 |
+
|
| 451 |
+
def _build_grid(beats: Sequence[BeatEvent], measures: Sequence[Measure]):
|
| 452 |
+
interval_denominators = np.zeros(len(beats) - 1, dtype=np.int32)
|
| 453 |
+
for measure in measures:
|
| 454 |
+
interval_denominators[measure.start_beat:measure.end_beat] = measure.denominator
|
| 455 |
+
if np.any(interval_denominators == 0):
|
| 456 |
+
raise BeatGridError("Downbeat spans do not cover every beat interval")
|
| 457 |
+
|
| 458 |
+
subbeat_times = []
|
| 459 |
+
subbeat_denominators = []
|
| 460 |
+
quarter_positions = [0.0]
|
| 461 |
+
current_quarter = 0.0
|
| 462 |
+
for index in range(len(beats) - 1):
|
| 463 |
+
start = beats[index].time
|
| 464 |
+
end = beats[index + 1].time
|
| 465 |
+
denominator = int(interval_denominators[index])
|
| 466 |
+
times = np.linspace(start, end, SUBBEAT_DIVISION + 1)[:-1]
|
| 467 |
+
subbeat_times.extend(float(value) for value in times)
|
| 468 |
+
subbeat_denominators.extend([denominator] * SUBBEAT_DIVISION)
|
| 469 |
+
quarter_step = 4.0 / denominator / SUBBEAT_DIVISION
|
| 470 |
+
for _ in range(SUBBEAT_DIVISION):
|
| 471 |
+
current_quarter += quarter_step
|
| 472 |
+
quarter_positions.append(current_quarter)
|
| 473 |
+
subbeat_times.append(beats[-1].time)
|
| 474 |
+
subbeat_denominators.append(int(interval_denominators[-1]))
|
| 475 |
+
return (
|
| 476 |
+
np.asarray(subbeat_times, dtype=np.float64),
|
| 477 |
+
np.asarray(quarter_positions, dtype=np.float64),
|
| 478 |
+
np.asarray(subbeat_denominators, dtype=np.int32),
|
| 479 |
+
)
|
| 480 |
+
|
| 481 |
+
|
| 482 |
+
def _subbeat_boundaries(subbeat_times: np.ndarray) -> np.ndarray:
|
| 483 |
+
return (subbeat_times[:-1] + subbeat_times[1:]) / 2
|
| 484 |
+
|
| 485 |
+
|
| 486 |
+
def _quantize_time(time: float, subbeat_times: np.ndarray) -> int:
|
| 487 |
+
return int(np.searchsorted(_subbeat_boundaries(subbeat_times), float(time)))
|
| 488 |
+
|
| 489 |
+
|
| 490 |
+
def _fill_intervals(rows, subbeat_times, *, default, dtype):
|
| 491 |
+
result = np.full(len(subbeat_times), default, dtype=dtype)
|
| 492 |
+
for start, end, value in rows:
|
| 493 |
+
start_t = _quantize_time(start, subbeat_times)
|
| 494 |
+
end_t = _quantize_time(end, subbeat_times)
|
| 495 |
+
start_t = max(0, min(start_t, len(result) - 1))
|
| 496 |
+
end_t = max(0, min(end_t, len(result) - 1))
|
| 497 |
+
if end_t <= start_t:
|
| 498 |
+
raise AbcRebuildError(
|
| 499 |
+
f"Interval {start:.6f}-{end:.6f} ({value}) is shorter than the ABC subbeat grid"
|
| 500 |
+
)
|
| 501 |
+
result[start_t:end_t] = value
|
| 502 |
+
if len(result) > 1:
|
| 503 |
+
result[-1] = result[-2]
|
| 504 |
+
return result
|
| 505 |
+
|
| 506 |
+
|
| 507 |
+
def _structure_events(rows, subbeat_times):
|
| 508 |
+
events = []
|
| 509 |
+
for start, _, label in rows:
|
| 510 |
+
t = _quantize_time(start, subbeat_times)
|
| 511 |
+
t = max(0, min(t, len(subbeat_times) - 1))
|
| 512 |
+
events.append((t, label))
|
| 513 |
+
return events
|
| 514 |
+
|
| 515 |
+
|
| 516 |
+
def _classify_melody_tracks(midi: pretty_midi.PrettyMIDI):
|
| 517 |
+
classified = {"Vocal": [], "Ins": []}
|
| 518 |
+
unknown = []
|
| 519 |
+
for instrument in midi.instruments:
|
| 520 |
+
if instrument.is_drum:
|
| 521 |
+
continue
|
| 522 |
+
name = (instrument.name or "").strip().lower()
|
| 523 |
+
if "vocal" in name:
|
| 524 |
+
classified["Vocal"].append(instrument)
|
| 525 |
+
elif "ins" in name or "instrument" in name:
|
| 526 |
+
classified["Ins"].append(instrument)
|
| 527 |
+
elif instrument.notes:
|
| 528 |
+
unknown.append(instrument)
|
| 529 |
+
if unknown:
|
| 530 |
+
if not classified["Vocal"] and not classified["Ins"] and len(unknown) == 1:
|
| 531 |
+
classified["Ins"].extend(unknown)
|
| 532 |
+
else:
|
| 533 |
+
names = [instrument.name or "<unnamed>" for instrument in unknown]
|
| 534 |
+
raise MelodyVoiceError(
|
| 535 |
+
f"Cannot map non-empty melody track(s) {names!r} to fixed Vocal/Ins voices"
|
| 536 |
+
)
|
| 537 |
+
return classified
|
| 538 |
+
|
| 539 |
+
|
| 540 |
+
def _notes_to_arr(notes, subbeat_times, voice_id):
|
| 541 |
+
result = np.zeros(len(subbeat_times), dtype=np.int32)
|
| 542 |
+
boundaries = _subbeat_boundaries(subbeat_times)
|
| 543 |
+
for note in sorted(notes, key=lambda item: (item.start, item.end, item.pitch)):
|
| 544 |
+
start_t = int(np.searchsorted(boundaries, note.start))
|
| 545 |
+
end_t = int(np.searchsorted(boundaries, note.end))
|
| 546 |
+
start_t = max(0, min(start_t, len(result) - 1))
|
| 547 |
+
end_t = max(0, min(end_t, len(result) - 1))
|
| 548 |
+
if end_t <= start_t:
|
| 549 |
+
raise MelodyVoiceError(
|
| 550 |
+
f"{voice_id}: MIDI note pitch={note.pitch} at {note.start:.6f}-{note.end:.6f} "
|
| 551 |
+
"cannot be represented on the decoded subbeat grid"
|
| 552 |
+
)
|
| 553 |
+
if np.any(result[start_t:end_t] != 0):
|
| 554 |
+
raise MelodyVoiceError(
|
| 555 |
+
f"{voice_id}: overlapping quantized melody notes at subbeats {start_t}:{end_t}"
|
| 556 |
+
)
|
| 557 |
+
sustain = note.pitch * 2 + 2
|
| 558 |
+
result[start_t:end_t] = sustain
|
| 559 |
+
result[start_t] = sustain + 1
|
| 560 |
+
return result
|
| 561 |
+
|
| 562 |
+
|
| 563 |
+
def _pitch_class(root: str) -> tuple[int, str, str]:
|
| 564 |
+
match = _ROOT_RE.fullmatch(root)
|
| 565 |
+
if match is None:
|
| 566 |
+
raise ChordSymbolError(f"Invalid pitch spelling {root!r}")
|
| 567 |
+
letter = match.group("letter")
|
| 568 |
+
accidental = match.group("accidental")
|
| 569 |
+
offset = accidental.count("#") - accidental.count("b")
|
| 570 |
+
return (_NATURAL_PITCH_CLASS[letter] + offset) % 12, letter, accidental
|
| 571 |
+
|
| 572 |
+
|
| 573 |
+
def portable_pitch_name(root: str, *, preserve_double: bool = False) -> str:
|
| 574 |
+
pitch_class, _, accidental = _pitch_class(root)
|
| 575 |
+
if preserve_double or len(accidental) <= 1:
|
| 576 |
+
return root
|
| 577 |
+
names = _SHARP_PITCH_NAMES if accidental.startswith("#") else _FLAT_PITCH_NAMES
|
| 578 |
+
return names[pitch_class]
|
| 579 |
+
|
| 580 |
+
|
| 581 |
+
def _bass_degree_to_pitch(root: str, degree_text: str) -> str:
|
| 582 |
+
if _ROOT_RE.fullmatch(degree_text):
|
| 583 |
+
return portable_pitch_name(degree_text, preserve_double=True)
|
| 584 |
+
match = _BASS_DEGREE_RE.fullmatch(degree_text)
|
| 585 |
+
if match is None:
|
| 586 |
+
raise ChordSymbolError(f"Invalid chord bass degree {degree_text!r}")
|
| 587 |
+
root_pc, root_letter, root_accidental = _pitch_class(root)
|
| 588 |
+
degree = int(match.group("degree"))
|
| 589 |
+
degree_accidental = match.group("accidental")
|
| 590 |
+
scale_semitones = (0, 2, 4, 5, 7, 9, 11)
|
| 591 |
+
interval = scale_semitones[(degree - 1) % 7] + 12 * ((degree - 1) // 7)
|
| 592 |
+
interval += degree_accidental.count("#") - degree_accidental.count("b")
|
| 593 |
+
target_pc = (root_pc + interval) % 12
|
| 594 |
+
|
| 595 |
+
target_letter_index = (_LETTERS.index(root_letter) + degree - 1) % 7
|
| 596 |
+
target_letter = _LETTERS[target_letter_index]
|
| 597 |
+
natural_pc = _NATURAL_PITCH_CLASS[target_letter]
|
| 598 |
+
difference = (target_pc - natural_pc + 6) % 12 - 6
|
| 599 |
+
if difference in {-2, -1, 0, 1, 2}:
|
| 600 |
+
accidental = {-2: "bb", -1: "b", 0: "", 1: "#", 2: "##"}[difference]
|
| 601 |
+
return target_letter + accidental
|
| 602 |
+
names = _SHARP_PITCH_NAMES if "#" in (root_accidental + degree_accidental) else _FLAT_PITCH_NAMES
|
| 603 |
+
return names[target_pc]
|
| 604 |
+
|
| 605 |
+
|
| 606 |
+
def chord_symbol_to_abc(chord: str) -> str | None:
|
| 607 |
+
chord = chord.strip()
|
| 608 |
+
if chord in NO_CHORDS:
|
| 609 |
+
return None
|
| 610 |
+
if ":" not in chord:
|
| 611 |
+
raise ChordSymbolError(f"Chord {chord!r} is missing the ':' quality separator")
|
| 612 |
+
root, descriptor = chord.split(":", 1)
|
| 613 |
+
if "/" in descriptor:
|
| 614 |
+
quality, bass_degree = descriptor.split("/", 1)
|
| 615 |
+
else:
|
| 616 |
+
quality, bass_degree = descriptor, None
|
| 617 |
+
if quality not in _QUALITY_TO_ABC:
|
| 618 |
+
raise ChordSymbolError(
|
| 619 |
+
f"Unsupported chord quality {quality!r} in {chord!r}; refusing to rewrite it as major"
|
| 620 |
+
)
|
| 621 |
+
chord_root = portable_pitch_name(root, preserve_double=True)
|
| 622 |
+
text = chord_root + _QUALITY_TO_ABC[quality]
|
| 623 |
+
if bass_degree:
|
| 624 |
+
text += "/" + _bass_degree_to_pitch(root, bass_degree)
|
| 625 |
+
return text
|
| 626 |
+
|
| 627 |
+
|
| 628 |
+
def key_symbol_to_abc(key: str) -> str:
|
| 629 |
+
key = key.strip()
|
| 630 |
+
if ":" in key:
|
| 631 |
+
root, mode = key.split(":", 1)
|
| 632 |
+
if mode not in {"major", "minor"}:
|
| 633 |
+
raise AbcRebuildError(f"Unsupported key mode {mode!r} in {key!r}")
|
| 634 |
+
elif key.endswith("m"):
|
| 635 |
+
root, mode = key[:-1], "minor"
|
| 636 |
+
else:
|
| 637 |
+
root, mode = key, "major"
|
| 638 |
+
root_pc, _, accidental = _pitch_class(root)
|
| 639 |
+
candidate = portable_pitch_name(root) + ("m" if mode == "minor" else "")
|
| 640 |
+
if candidate in _KEY_SIGNATURE_ACCIDENTALS:
|
| 641 |
+
return candidate
|
| 642 |
+
names = _FLAT_PITCH_NAMES if "b" in accidental else _SHARP_PITCH_NAMES
|
| 643 |
+
candidate = names[root_pc] + ("m" if mode == "minor" else "")
|
| 644 |
+
if candidate not in _KEY_SIGNATURE_ACCIDENTALS:
|
| 645 |
+
fallback_names = _SHARP_PITCH_NAMES if names is _FLAT_PITCH_NAMES else _FLAT_PITCH_NAMES
|
| 646 |
+
candidate = fallback_names[root_pc] + ("m" if mode == "minor" else "")
|
| 647 |
+
if candidate not in _KEY_SIGNATURE_ACCIDENTALS:
|
| 648 |
+
raise AbcRebuildError(f"Cannot encode portable ABC key for {key!r}")
|
| 649 |
+
return candidate
|
| 650 |
+
|
| 651 |
+
|
| 652 |
+
def get_key_accidentals(key: str) -> list[int]:
|
| 653 |
+
try:
|
| 654 |
+
count = _KEY_SIGNATURE_ACCIDENTALS[key]
|
| 655 |
+
except KeyError as exc:
|
| 656 |
+
raise AbcRebuildError(f"Unsupported ABC key signature {key!r}") from exc
|
| 657 |
+
accidentals = [0] * 7
|
| 658 |
+
order = "FCGDAEB" if count > 0 else "BEADGCF"
|
| 659 |
+
for letter in order[:abs(count)]:
|
| 660 |
+
accidentals[_LETTERS.index(letter)] = 1 if count > 0 else -1
|
| 661 |
+
return accidentals
|
| 662 |
+
|
| 663 |
+
|
| 664 |
+
def note_to_abc(note: int, key_accidentals: Sequence[int], measure_accidentals: dict) -> str:
|
| 665 |
+
"""Use key-relative spelling and write only bar-state changes.
|
| 666 |
+
|
| 667 |
+
The two target parsers propagate an accidental to the same note letter in
|
| 668 |
+
every octave until the next barline. ``measure_accidentals`` is therefore
|
| 669 |
+
keyed by letter and reset by the caller for every bar (and after an inline
|
| 670 |
+
key change). This preserves pitches across parsers while still omitting
|
| 671 |
+
repeated accidental marks. The key-relative spelling can use double
|
| 672 |
+
accidentals in remote keys; MIDI G is F## in G# minor, for example.
|
| 673 |
+
"""
|
| 674 |
+
|
| 675 |
+
accidental_count = sum(key_accidentals)
|
| 676 |
+
try:
|
| 677 |
+
pitch_name = _KEY_RELATIVE_PITCH_NAMES[accidental_count][note % 12]
|
| 678 |
+
except KeyError as exc:
|
| 679 |
+
raise AbcRebuildError(
|
| 680 |
+
f"Unsupported key signature accidental count {accidental_count}"
|
| 681 |
+
) from exc
|
| 682 |
+
letter = pitch_name[0]
|
| 683 |
+
accidental = pitch_name[1:]
|
| 684 |
+
accidental_number = {"": 0, "#": 1, "##": 2, "b": -1, "bb": -2}[accidental]
|
| 685 |
+
octave = (note - 60) // 12
|
| 686 |
+
# Cb and B# cross the MIDI octave boundary even though their written note
|
| 687 |
+
# letter does not.
|
| 688 |
+
if note % 12 == 11 and accidental_number == -1:
|
| 689 |
+
octave += 1
|
| 690 |
+
elif note % 12 == 0 and accidental_number == 1:
|
| 691 |
+
octave -= 1
|
| 692 |
+
scale_index = _LETTERS.index(letter)
|
| 693 |
+
current_accidental = measure_accidentals.get(
|
| 694 |
+
scale_index,
|
| 695 |
+
key_accidentals[scale_index],
|
| 696 |
+
)
|
| 697 |
+
accidental_text = ""
|
| 698 |
+
if current_accidental != accidental_number:
|
| 699 |
+
measure_accidentals[scale_index] = accidental_number
|
| 700 |
+
accidental_text = {-2: "__", -1: "_", 0: "=", 1: "^", 2: "^^"}[
|
| 701 |
+
accidental_number
|
| 702 |
+
]
|
| 703 |
+
|
| 704 |
+
if octave > 0:
|
| 705 |
+
letter = letter.lower()
|
| 706 |
+
if octave > 1:
|
| 707 |
+
letter += "'" * (octave - 1)
|
| 708 |
+
elif octave < 0:
|
| 709 |
+
letter += "," * abs(octave)
|
| 710 |
+
return accidental_text + letter
|
| 711 |
+
|
| 712 |
+
|
| 713 |
+
def build_rebuilt_abc_score(
|
| 714 |
+
melody_midi_path,
|
| 715 |
+
beats_path,
|
| 716 |
+
chords_path,
|
| 717 |
+
keys_path,
|
| 718 |
+
structures_path,
|
| 719 |
+
*,
|
| 720 |
+
meter_conflict: str = "infer",
|
| 721 |
+
melody_only: bool = False,
|
| 722 |
+
) -> RebuiltAbcScore:
|
| 723 |
+
beats = read_beats(beats_path)
|
| 724 |
+
keys = read_keys(keys_path)
|
| 725 |
+
structures = read_structures(structures_path)
|
| 726 |
+
chords = [] if melody_only else read_chords(chords_path)
|
| 727 |
+
midi = pretty_midi.PrettyMIDI(str(melody_midi_path))
|
| 728 |
+
return _assemble_abc_score(midi, beats, keys, structures, chords,
|
| 729 |
+
meter_conflict=meter_conflict, melody_only=melody_only)
|
| 730 |
+
|
| 731 |
+
|
| 732 |
+
def build_rebuilt_abc_score_from_data(
|
| 733 |
+
melody_midi, beats, chords, keys, structures, *, meter_conflict="infer", melody_only=False,
|
| 734 |
+
) -> RebuiltAbcScore:
|
| 735 |
+
"""Build from MIDI bytes/BytesIO/PrettyMIDI and beat/interval rows.
|
| 736 |
+
|
| 737 |
+
BeatEvent lists are also accepted. The file and memory interfaces share
|
| 738 |
+
interval validation, score construction, serialization and ABC validation.
|
| 739 |
+
"""
|
| 740 |
+
beats = list(beats)
|
| 741 |
+
if beats and isinstance(beats[0], BeatEvent):
|
| 742 |
+
beat_entries = [(b.line_no, [str(b.time), str(b.beat_id), str(b.declared_numerator), str(b.denominator)]) for b in beats]
|
| 743 |
+
else:
|
| 744 |
+
beat_entries = _row_entries(beats, "beats")
|
| 745 |
+
beats = _parse_beats(beat_entries, "beats")
|
| 746 |
+
keys = _parse_keys(_row_entries(keys, "keys"), "keys")
|
| 747 |
+
structures = _parse_structures(_row_entries(structures, "structures"), "structures")
|
| 748 |
+
chords = [] if melody_only else _parse_chords(_row_entries(chords, "chords"), "chords")
|
| 749 |
+
if isinstance(melody_midi, (bytes, bytearray)):
|
| 750 |
+
melody_midi = BytesIO(melody_midi)
|
| 751 |
+
if not isinstance(melody_midi, (BytesIO, pretty_midi.PrettyMIDI)):
|
| 752 |
+
raise TypeError("melody_midi must be MIDI bytes, BytesIO, or PrettyMIDI")
|
| 753 |
+
midi = melody_midi if isinstance(melody_midi, pretty_midi.PrettyMIDI) else pretty_midi.PrettyMIDI(melody_midi)
|
| 754 |
+
return _assemble_abc_score(midi, beats, keys, structures, chords,
|
| 755 |
+
meter_conflict=meter_conflict, melody_only=melody_only)
|
| 756 |
+
|
| 757 |
+
|
| 758 |
+
def _assemble_abc_score(midi, beats, keys, structures, chords, *, meter_conflict, melody_only):
|
| 759 |
+
measures, diagnostics = infer_measures(beats, meter_conflict=meter_conflict)
|
| 760 |
+
subbeat_times, subbeat_quarters, subbeat_denominators = _build_grid(beats, measures)
|
| 761 |
+
classified = _classify_melody_tracks(midi)
|
| 762 |
+
voice_arrs = {}
|
| 763 |
+
for voice_id in VOICE_IDS:
|
| 764 |
+
notes = [
|
| 765 |
+
note
|
| 766 |
+
for instrument in classified[voice_id]
|
| 767 |
+
for note in instrument.notes
|
| 768 |
+
]
|
| 769 |
+
voice_arrs[voice_id] = _notes_to_arr(notes, subbeat_times, voice_id)
|
| 770 |
+
|
| 771 |
+
key_arr = _fill_intervals(keys, subbeat_times, default=keys[0][2], dtype="<U16")
|
| 772 |
+
if melody_only:
|
| 773 |
+
# Do not even read chord labels in melody-only mode. A constant no-chord
|
| 774 |
+
# timeline removes chord-only render boundaries, allowing held notes and
|
| 775 |
+
# rests to be serialized as their original semantic segments.
|
| 776 |
+
chord_arr = np.full(len(subbeat_times), "N", dtype="<U64")
|
| 777 |
+
else:
|
| 778 |
+
chord_arr = _fill_intervals(
|
| 779 |
+
chords,
|
| 780 |
+
subbeat_times,
|
| 781 |
+
default="N",
|
| 782 |
+
dtype="<U64",
|
| 783 |
+
)
|
| 784 |
+
return RebuiltAbcScore(
|
| 785 |
+
beats=list(beats),
|
| 786 |
+
measures=measures,
|
| 787 |
+
subbeat_times=subbeat_times,
|
| 788 |
+
subbeat_quarters=subbeat_quarters,
|
| 789 |
+
subbeat_denominators=subbeat_denominators,
|
| 790 |
+
key_arr=key_arr,
|
| 791 |
+
chord_arr=chord_arr,
|
| 792 |
+
structure_events=_structure_events(structures, subbeat_times),
|
| 793 |
+
voice_arrs=voice_arrs,
|
| 794 |
+
diagnostics=diagnostics,
|
| 795 |
+
)
|
| 796 |
+
|
| 797 |
+
|
| 798 |
+
def abc_unit_denominator(score: RebuiltAbcScore) -> int:
|
| 799 |
+
values = [
|
| 800 |
+
denominator * score.subbeat_div
|
| 801 |
+
for measure in score.measures
|
| 802 |
+
for denominator in (measure.denominator, measure.abc_denominator)
|
| 803 |
+
]
|
| 804 |
+
denominator = math.lcm(*values)
|
| 805 |
+
if denominator > 1024:
|
| 806 |
+
raise AbcRebuildError(f"Required ABC unit length 1/{denominator} is unreasonably small")
|
| 807 |
+
return denominator
|
| 808 |
+
|
| 809 |
+
|
| 810 |
+
def _measure_actual_units(measure: Measure, unit_denominator: int) -> int:
|
| 811 |
+
return measure.numerator * unit_denominator // measure.denominator
|
| 812 |
+
|
| 813 |
+
|
| 814 |
+
def _measure_abc_units(measure: Measure, unit_denominator: int) -> int:
|
| 815 |
+
return measure.abc_numerator * unit_denominator // measure.abc_denominator
|
| 816 |
+
|
| 817 |
+
|
| 818 |
+
def _measure_padding_units(measure: Measure, unit_denominator: int) -> int:
|
| 819 |
+
return (
|
| 820 |
+
_measure_abc_units(measure, unit_denominator)
|
| 821 |
+
- _measure_actual_units(measure, unit_denominator)
|
| 822 |
+
)
|
| 823 |
+
|
| 824 |
+
|
| 825 |
+
def _duration_units(score: RebuiltAbcScore, start_t: int, end_t: int, unit_denominator: int) -> int:
|
| 826 |
+
units = 0
|
| 827 |
+
for denominator in score.subbeat_denominators[start_t:end_t]:
|
| 828 |
+
divisor = int(denominator) * score.subbeat_div
|
| 829 |
+
if unit_denominator % divisor:
|
| 830 |
+
raise AbcRebuildError(
|
| 831 |
+
f"ABC L:1/{unit_denominator} cannot express a 1/{divisor} subbeat exactly"
|
| 832 |
+
)
|
| 833 |
+
units += unit_denominator // divisor
|
| 834 |
+
return units
|
| 835 |
+
|
| 836 |
+
|
| 837 |
+
def estimate_tempo(score: RebuiltAbcScore) -> float:
|
| 838 |
+
seconds = score.subbeat_times[-1] - score.subbeat_times[0]
|
| 839 |
+
quarter_notes = score.subbeat_quarters[-1] - score.subbeat_quarters[0]
|
| 840 |
+
if seconds <= 0 or quarter_notes <= 0:
|
| 841 |
+
raise AbcRebuildError("Cannot estimate tempo from a zero-duration score")
|
| 842 |
+
return float(quarter_notes / seconds * 60.0)
|
| 843 |
+
|
| 844 |
+
|
| 845 |
+
def _continues_pitch(value: int, next_value: int) -> bool:
|
| 846 |
+
if value <= 0:
|
| 847 |
+
return False
|
| 848 |
+
pitch = value // 2 - 1
|
| 849 |
+
return next_value == pitch * 2 + 2
|
| 850 |
+
|
| 851 |
+
|
| 852 |
+
def _same_note_segment(value: int, next_value: int) -> bool:
|
| 853 |
+
if value == 0:
|
| 854 |
+
return next_value == 0
|
| 855 |
+
pitch = value // 2 - 1
|
| 856 |
+
return next_value == pitch * 2 + 2
|
| 857 |
+
|
| 858 |
+
|
| 859 |
+
_SUPPORTED_DURATION_UNITS = frozenset(
|
| 860 |
+
{1, 2, 3, 4, 6, 8, 12, 16, 24, 32, 48}
|
| 861 |
+
)
|
| 862 |
+
|
| 863 |
+
|
| 864 |
+
def _split_duration_units(duration: int) -> list[int]:
|
| 865 |
+
"""Split a duration into values accepted by strict music parsers."""
|
| 866 |
+
|
| 867 |
+
if duration <= 0:
|
| 868 |
+
raise AbcRebuildError(f"Cannot serialize non-positive duration {duration}")
|
| 869 |
+
result = []
|
| 870 |
+
remaining = int(duration)
|
| 871 |
+
while remaining:
|
| 872 |
+
if remaining in _SUPPORTED_DURATION_UNITS:
|
| 873 |
+
result.append(remaining)
|
| 874 |
+
break
|
| 875 |
+
candidates = [
|
| 876 |
+
value
|
| 877 |
+
for value in _SUPPORTED_DURATION_UNITS
|
| 878 |
+
if value < remaining
|
| 879 |
+
]
|
| 880 |
+
if not candidates:
|
| 881 |
+
raise AbcRebuildError(
|
| 882 |
+
f"Duration {duration} cannot be split into representable ABC values"
|
| 883 |
+
)
|
| 884 |
+
chunk = max(candidates)
|
| 885 |
+
result.append(chunk)
|
| 886 |
+
remaining -= chunk
|
| 887 |
+
return result
|
| 888 |
+
|
| 889 |
+
|
| 890 |
+
def _duration_text(duration: int) -> str:
|
| 891 |
+
return "" if duration == 1 else str(duration)
|
| 892 |
+
|
| 893 |
+
|
| 894 |
+
def _render_duration_tokens(
|
| 895 |
+
prefix: str,
|
| 896 |
+
note_text: str,
|
| 897 |
+
duration: int,
|
| 898 |
+
*,
|
| 899 |
+
tie_out: bool,
|
| 900 |
+
) -> list[str]:
|
| 901 |
+
chunks = _split_duration_units(duration)
|
| 902 |
+
tokens = []
|
| 903 |
+
for index, chunk in enumerate(chunks):
|
| 904 |
+
continues = note_text != "z" and (
|
| 905 |
+
index + 1 < len(chunks) or tie_out
|
| 906 |
+
)
|
| 907 |
+
tokens.append(
|
| 908 |
+
(prefix if index == 0 else "")
|
| 909 |
+
+ note_text
|
| 910 |
+
+ _duration_text(chunk)
|
| 911 |
+
+ ("-" if continues else "")
|
| 912 |
+
)
|
| 913 |
+
return tokens
|
| 914 |
+
|
| 915 |
+
|
| 916 |
+
def _render_voice_measure(
|
| 917 |
+
score: RebuiltAbcScore,
|
| 918 |
+
voice_id: str,
|
| 919 |
+
measure: Measure,
|
| 920 |
+
unit_denominator: int,
|
| 921 |
+
) -> str:
|
| 922 |
+
voice = score.voice_arrs[voice_id]
|
| 923 |
+
show_chords = voice_id == "Vocal"
|
| 924 |
+
measure_accidentals = {}
|
| 925 |
+
current_key = str(score.key_arr[measure.start_t])
|
| 926 |
+
key_accidentals = get_key_accidentals(current_key)
|
| 927 |
+
parts = []
|
| 928 |
+
padding = _measure_padding_units(measure, unit_denominator)
|
| 929 |
+
if padding < 0:
|
| 930 |
+
raise AbcRebuildError(
|
| 931 |
+
f"Measure {measure.index}: notated meter is shorter than its decoded span"
|
| 932 |
+
)
|
| 933 |
+
leading_padding = padding if measure.pad_before else 0
|
| 934 |
+
trailing_padding = 0 if measure.pad_before else padding
|
| 935 |
+
t = measure.start_t
|
| 936 |
+
while t < measure.end_t:
|
| 937 |
+
change_points = [measure.end_t]
|
| 938 |
+
for probe in range(t + 1, measure.end_t):
|
| 939 |
+
if not _same_note_segment(int(voice[t]), int(voice[probe])):
|
| 940 |
+
change_points.append(probe)
|
| 941 |
+
break
|
| 942 |
+
for probe in range(t + 1, measure.end_t):
|
| 943 |
+
if score.key_arr[probe] != score.key_arr[probe - 1]:
|
| 944 |
+
change_points.append(probe)
|
| 945 |
+
break
|
| 946 |
+
if show_chords:
|
| 947 |
+
for probe in range(t + 1, measure.end_t):
|
| 948 |
+
if score.chord_arr[probe] != score.chord_arr[probe - 1]:
|
| 949 |
+
change_points.append(probe)
|
| 950 |
+
break
|
| 951 |
+
next_t = min(change_points)
|
| 952 |
+
|
| 953 |
+
prefix = ""
|
| 954 |
+
key = str(score.key_arr[t])
|
| 955 |
+
if t > measure.start_t and key != current_key:
|
| 956 |
+
current_key = key
|
| 957 |
+
key_accidentals = get_key_accidentals(current_key)
|
| 958 |
+
measure_accidentals = {}
|
| 959 |
+
prefix += f"[K:{current_key}]"
|
| 960 |
+
|
| 961 |
+
if show_chords and (t == measure.start_t or score.chord_arr[t] != score.chord_arr[t - 1]):
|
| 962 |
+
chord = str(score.chord_arr[t])
|
| 963 |
+
chord_text = chord_symbol_to_abc(chord)
|
| 964 |
+
if chord_text is not None:
|
| 965 |
+
prefix += f'"{chord_text}"'
|
| 966 |
+
|
| 967 |
+
value = int(voice[t])
|
| 968 |
+
if value == 0:
|
| 969 |
+
note_text = "z"
|
| 970 |
+
else:
|
| 971 |
+
note_text = note_to_abc(value // 2 - 1, key_accidentals, measure_accidentals)
|
| 972 |
+
duration = _duration_units(score, t, next_t, unit_denominator)
|
| 973 |
+
if t == measure.start_t and leading_padding:
|
| 974 |
+
if value == 0 and not prefix:
|
| 975 |
+
duration += leading_padding
|
| 976 |
+
else:
|
| 977 |
+
parts.extend(
|
| 978 |
+
_render_duration_tokens(
|
| 979 |
+
"",
|
| 980 |
+
"z",
|
| 981 |
+
leading_padding,
|
| 982 |
+
tie_out=False,
|
| 983 |
+
)
|
| 984 |
+
)
|
| 985 |
+
leading_padding = 0
|
| 986 |
+
if value == 0 and next_t == measure.end_t and trailing_padding:
|
| 987 |
+
duration += trailing_padding
|
| 988 |
+
trailing_padding = 0
|
| 989 |
+
if duration <= 0:
|
| 990 |
+
raise AbcRebuildError(f"Non-positive ABC duration at subbeats {t}:{next_t}")
|
| 991 |
+
tie_out = (
|
| 992 |
+
value > 0
|
| 993 |
+
and next_t < len(voice)
|
| 994 |
+
and _continues_pitch(value, int(voice[next_t]))
|
| 995 |
+
)
|
| 996 |
+
parts.extend(
|
| 997 |
+
_render_duration_tokens(
|
| 998 |
+
prefix,
|
| 999 |
+
note_text,
|
| 1000 |
+
duration,
|
| 1001 |
+
tie_out=tie_out,
|
| 1002 |
+
)
|
| 1003 |
+
)
|
| 1004 |
+
t = next_t
|
| 1005 |
+
if leading_padding:
|
| 1006 |
+
raise AbcRebuildError(
|
| 1007 |
+
f"Measure {measure.index}: leading rest padding was not serialized"
|
| 1008 |
+
)
|
| 1009 |
+
if trailing_padding:
|
| 1010 |
+
parts.extend(
|
| 1011 |
+
_render_duration_tokens(
|
| 1012 |
+
"",
|
| 1013 |
+
"z",
|
| 1014 |
+
trailing_padding,
|
| 1015 |
+
tie_out=False,
|
| 1016 |
+
)
|
| 1017 |
+
)
|
| 1018 |
+
return "".join(parts)
|
| 1019 |
+
|
| 1020 |
+
|
| 1021 |
+
def _is_compressible_full_rest(rendered_measure: str) -> bool:
|
| 1022 |
+
"""Whether a rendered measure can be losslessly replaced by ABC ``Z``."""
|
| 1023 |
+
|
| 1024 |
+
cursor = 0
|
| 1025 |
+
saw_note = False
|
| 1026 |
+
for match in _MUSIC_ELEMENT_RE.finditer(rendered_measure):
|
| 1027 |
+
if rendered_measure[cursor:match.start()]:
|
| 1028 |
+
return False
|
| 1029 |
+
cursor = match.end()
|
| 1030 |
+
if match.group("quoted") is not None or match.group("key") is not None:
|
| 1031 |
+
return False
|
| 1032 |
+
saw_note = True
|
| 1033 |
+
if match.group("note") != "z" or match.group("tie"):
|
| 1034 |
+
return False
|
| 1035 |
+
return saw_note and cursor == len(rendered_measure)
|
| 1036 |
+
|
| 1037 |
+
|
| 1038 |
+
def _render_voice_group(
|
| 1039 |
+
score: RebuiltAbcScore,
|
| 1040 |
+
voice_id: str,
|
| 1041 |
+
measures: list[Measure],
|
| 1042 |
+
unit_denominator: int,
|
| 1043 |
+
) -> str:
|
| 1044 |
+
rendered = [
|
| 1045 |
+
_render_voice_measure(
|
| 1046 |
+
score,
|
| 1047 |
+
voice_id,
|
| 1048 |
+
measure,
|
| 1049 |
+
unit_denominator,
|
| 1050 |
+
)
|
| 1051 |
+
for measure in measures
|
| 1052 |
+
]
|
| 1053 |
+
parts = []
|
| 1054 |
+
index = 0
|
| 1055 |
+
while index < len(rendered):
|
| 1056 |
+
if not _is_compressible_full_rest(rendered[index]):
|
| 1057 |
+
parts.append(rendered[index] + "|")
|
| 1058 |
+
index += 1
|
| 1059 |
+
continue
|
| 1060 |
+
end = index + 1
|
| 1061 |
+
while (
|
| 1062 |
+
end < len(rendered)
|
| 1063 |
+
and _is_compressible_full_rest(rendered[end])
|
| 1064 |
+
):
|
| 1065 |
+
end += 1
|
| 1066 |
+
count = end - index
|
| 1067 |
+
parts.append("Z" + (str(count) if count > 1 else "") + "|")
|
| 1068 |
+
index = end
|
| 1069 |
+
return "".join(parts)
|
| 1070 |
+
|
| 1071 |
+
|
| 1072 |
+
def _sanitize_structure_label(value: str) -> str:
|
| 1073 |
+
return " ".join(str(value).split())
|
| 1074 |
+
|
| 1075 |
+
|
| 1076 |
+
def _measure_groups(score: RebuiltAbcScore) -> list[MeasureGroup]:
|
| 1077 |
+
first_measure = score.measures[0]
|
| 1078 |
+
active_meter = (
|
| 1079 |
+
first_measure.abc_numerator,
|
| 1080 |
+
first_measure.abc_denominator,
|
| 1081 |
+
)
|
| 1082 |
+
active_key = str(score.key_arr[first_measure.start_t])
|
| 1083 |
+
active_structure = ""
|
| 1084 |
+
groups: list[MeasureGroup] = []
|
| 1085 |
+
|
| 1086 |
+
for measure in score.measures:
|
| 1087 |
+
meter = (measure.abc_numerator, measure.abc_denominator)
|
| 1088 |
+
key = str(score.key_arr[measure.start_t])
|
| 1089 |
+
meter_changed = meter != active_meter
|
| 1090 |
+
key_changed = key != active_key
|
| 1091 |
+
new_structure_labels = []
|
| 1092 |
+
for t, label in score.structure_events:
|
| 1093 |
+
if not measure.start_t <= t < measure.end_t:
|
| 1094 |
+
continue
|
| 1095 |
+
clean_label = _sanitize_structure_label(label)
|
| 1096 |
+
if clean_label and clean_label != active_structure:
|
| 1097 |
+
new_structure_labels.append(clean_label)
|
| 1098 |
+
active_structure = clean_label
|
| 1099 |
+
|
| 1100 |
+
start_group = (
|
| 1101 |
+
not groups
|
| 1102 |
+
or len(groups[-1].measures) >= 4
|
| 1103 |
+
or meter_changed
|
| 1104 |
+
or key_changed
|
| 1105 |
+
or bool(new_structure_labels)
|
| 1106 |
+
)
|
| 1107 |
+
if start_group:
|
| 1108 |
+
groups.append(
|
| 1109 |
+
MeasureGroup(
|
| 1110 |
+
measures=[measure],
|
| 1111 |
+
structure_labels=new_structure_labels,
|
| 1112 |
+
meter_changed=meter_changed,
|
| 1113 |
+
key_changed=key_changed,
|
| 1114 |
+
)
|
| 1115 |
+
)
|
| 1116 |
+
else:
|
| 1117 |
+
groups[-1].measures.append(measure)
|
| 1118 |
+
|
| 1119 |
+
active_meter = meter
|
| 1120 |
+
active_key = str(score.key_arr[measure.end_t - 1])
|
| 1121 |
+
return groups
|
| 1122 |
+
|
| 1123 |
+
|
| 1124 |
+
def score_to_abc(score: RebuiltAbcScore) -> str:
|
| 1125 |
+
unit_denominator = abc_unit_denominator(score)
|
| 1126 |
+
first_measure = score.measures[0]
|
| 1127 |
+
first_key = str(score.key_arr[first_measure.start_t])
|
| 1128 |
+
lines = [
|
| 1129 |
+
"X:1",
|
| 1130 |
+
"T:",
|
| 1131 |
+
f"M:{first_measure.abc_numerator}/{first_measure.abc_denominator}",
|
| 1132 |
+
f"L:1/{unit_denominator}",
|
| 1133 |
+
f"Q:1/4={int(round(estimate_tempo(score)))}",
|
| 1134 |
+
'V: Vocal clef=treble name="Vocal Melody" snm="Vocal"',
|
| 1135 |
+
'V: Ins clef=treble name="Ins Melody" snm="Inst."',
|
| 1136 |
+
f"K:{first_key}",
|
| 1137 |
+
]
|
| 1138 |
+
for group in _measure_groups(score):
|
| 1139 |
+
lines.extend(f"% {label}" for label in group.structure_labels)
|
| 1140 |
+
first_group_measure = group.measures[0]
|
| 1141 |
+
for voice_id in VOICE_IDS:
|
| 1142 |
+
lines.append(f"V: {voice_id}")
|
| 1143 |
+
if group.meter_changed:
|
| 1144 |
+
lines.append(
|
| 1145 |
+
f"M:{first_group_measure.abc_numerator}/"
|
| 1146 |
+
f"{first_group_measure.abc_denominator}"
|
| 1147 |
+
)
|
| 1148 |
+
if group.key_changed:
|
| 1149 |
+
lines.append(
|
| 1150 |
+
f"K:{score.key_arr[first_group_measure.start_t]}"
|
| 1151 |
+
)
|
| 1152 |
+
lines.append(
|
| 1153 |
+
_render_voice_group(
|
| 1154 |
+
score,
|
| 1155 |
+
voice_id,
|
| 1156 |
+
group.measures,
|
| 1157 |
+
unit_denominator,
|
| 1158 |
+
)
|
| 1159 |
+
)
|
| 1160 |
+
text = "\n".join(lines) + "\n"
|
| 1161 |
+
validate_serialized_abc(text, score)
|
| 1162 |
+
return text
|
| 1163 |
+
|
| 1164 |
+
|
| 1165 |
+
_MUSIC_ELEMENT_RE = re.compile(
|
| 1166 |
+
r'"(?P<quoted>[^"]*)"'
|
| 1167 |
+
r"|\[K:(?P<key>[^\]]+)\]"
|
| 1168 |
+
r"|(?P<note>[_=^]*[A-Ga-gz][,']*)(?P<duration>\d*)(?P<tie>-?)"
|
| 1169 |
+
)
|
| 1170 |
+
|
| 1171 |
+
|
| 1172 |
+
def _parse_music_measure(line: str, expected: int, context: str):
|
| 1173 |
+
body = line
|
| 1174 |
+
if body == "Z":
|
| 1175 |
+
return [], []
|
| 1176 |
+
position = 0
|
| 1177 |
+
cursor = 0
|
| 1178 |
+
quoted_events = []
|
| 1179 |
+
key_events = []
|
| 1180 |
+
for match in _MUSIC_ELEMENT_RE.finditer(body):
|
| 1181 |
+
gap = body[cursor:match.start()]
|
| 1182 |
+
if gap.strip():
|
| 1183 |
+
raise AbcRebuildError(f"{context}: unsupported serialized ABC tokens {gap!r}")
|
| 1184 |
+
cursor = match.end()
|
| 1185 |
+
if match.group("quoted") is not None:
|
| 1186 |
+
quoted_events.append((position, match.group("quoted")))
|
| 1187 |
+
continue
|
| 1188 |
+
if match.group("key") is not None:
|
| 1189 |
+
key_events.append((position, match.group("key")))
|
| 1190 |
+
continue
|
| 1191 |
+
note = match.group("note")
|
| 1192 |
+
tie = match.group("tie")
|
| 1193 |
+
if tie and note == "z":
|
| 1194 |
+
raise AbcRebuildError(f"{context}: a rest cannot be tied")
|
| 1195 |
+
duration_text = match.group("duration")
|
| 1196 |
+
duration = int(duration_text) if duration_text else 1
|
| 1197 |
+
if duration not in _SUPPORTED_DURATION_UNITS:
|
| 1198 |
+
raise AbcRebuildError(
|
| 1199 |
+
f"{context}: duration {duration} is not parser-representable"
|
| 1200 |
+
)
|
| 1201 |
+
position += duration
|
| 1202 |
+
if body[cursor:].strip():
|
| 1203 |
+
raise AbcRebuildError(
|
| 1204 |
+
f"{context}: unsupported serialized ABC tokens {body[cursor:]!r}"
|
| 1205 |
+
)
|
| 1206 |
+
if position != expected:
|
| 1207 |
+
raise AbcRebuildError(
|
| 1208 |
+
f"{context}: duration {position} does not match meter duration {expected}"
|
| 1209 |
+
)
|
| 1210 |
+
if re.search(r"(^|[\s|])-[_=^A-Ga-g]", line):
|
| 1211 |
+
raise AbcRebuildError(f"{context}: tie is written before its second note")
|
| 1212 |
+
return quoted_events, key_events
|
| 1213 |
+
|
| 1214 |
+
|
| 1215 |
+
def _expected_measure_chords(
|
| 1216 |
+
score: RebuiltAbcScore,
|
| 1217 |
+
measure: Measure,
|
| 1218 |
+
unit_denominator: int,
|
| 1219 |
+
) -> list[tuple[int, str]]:
|
| 1220 |
+
expected = []
|
| 1221 |
+
leading_padding = (
|
| 1222 |
+
_measure_padding_units(measure, unit_denominator)
|
| 1223 |
+
if measure.pad_before
|
| 1224 |
+
else 0
|
| 1225 |
+
)
|
| 1226 |
+
for t in range(measure.start_t, measure.end_t):
|
| 1227 |
+
if t != measure.start_t and score.chord_arr[t] == score.chord_arr[t - 1]:
|
| 1228 |
+
continue
|
| 1229 |
+
position = leading_padding + _duration_units(
|
| 1230 |
+
score,
|
| 1231 |
+
measure.start_t,
|
| 1232 |
+
t,
|
| 1233 |
+
unit_denominator,
|
| 1234 |
+
)
|
| 1235 |
+
chord = str(score.chord_arr[t])
|
| 1236 |
+
text = chord_symbol_to_abc(chord)
|
| 1237 |
+
if text is not None:
|
| 1238 |
+
expected.append((position, text))
|
| 1239 |
+
return expected
|
| 1240 |
+
|
| 1241 |
+
|
| 1242 |
+
def _expected_measure_keys(
|
| 1243 |
+
score: RebuiltAbcScore,
|
| 1244 |
+
measure: Measure,
|
| 1245 |
+
unit_denominator: int,
|
| 1246 |
+
) -> list[tuple[int, str]]:
|
| 1247 |
+
leading_padding = (
|
| 1248 |
+
_measure_padding_units(measure, unit_denominator)
|
| 1249 |
+
if measure.pad_before
|
| 1250 |
+
else 0
|
| 1251 |
+
)
|
| 1252 |
+
return [
|
| 1253 |
+
(
|
| 1254 |
+
leading_padding
|
| 1255 |
+
+ _duration_units(score, measure.start_t, t, unit_denominator),
|
| 1256 |
+
str(score.key_arr[t]),
|
| 1257 |
+
)
|
| 1258 |
+
for t in range(measure.start_t + 1, measure.end_t)
|
| 1259 |
+
if score.key_arr[t] != score.key_arr[t - 1]
|
| 1260 |
+
]
|
| 1261 |
+
|
| 1262 |
+
|
| 1263 |
+
def _parse_voice_group(lines, cursor, voice_id, group_index):
|
| 1264 |
+
expected_voice_field = f"V: {voice_id}"
|
| 1265 |
+
if cursor >= len(lines) or lines[cursor] != expected_voice_field:
|
| 1266 |
+
observed = lines[cursor] if cursor < len(lines) else "<end>"
|
| 1267 |
+
raise AbcRebuildError(
|
| 1268 |
+
f"Group {group_index}: expected {expected_voice_field}, got {observed!r}"
|
| 1269 |
+
)
|
| 1270 |
+
cursor += 1
|
| 1271 |
+
fields = {}
|
| 1272 |
+
while cursor < len(lines) and (
|
| 1273 |
+
lines[cursor].startswith("M:")
|
| 1274 |
+
or lines[cursor].startswith("K:")
|
| 1275 |
+
):
|
| 1276 |
+
name, value = lines[cursor].split(":", 1)
|
| 1277 |
+
if name in fields:
|
| 1278 |
+
raise AbcRebuildError(
|
| 1279 |
+
f"Group {group_index} {voice_id}: repeated {name}: field"
|
| 1280 |
+
)
|
| 1281 |
+
fields[name] = value
|
| 1282 |
+
cursor += 1
|
| 1283 |
+
if cursor >= len(lines):
|
| 1284 |
+
raise AbcRebuildError(
|
| 1285 |
+
f"Group {group_index} {voice_id}: missing music line"
|
| 1286 |
+
)
|
| 1287 |
+
music_line = lines[cursor]
|
| 1288 |
+
if music_line.startswith(("V:", "M:", "K:", "%")):
|
| 1289 |
+
raise AbcRebuildError(
|
| 1290 |
+
f"Group {group_index} {voice_id}: invalid music line {music_line!r}"
|
| 1291 |
+
)
|
| 1292 |
+
cursor += 1
|
| 1293 |
+
split_bars = music_line.split("|")
|
| 1294 |
+
if not split_bars or split_bars[-1].strip():
|
| 1295 |
+
raise AbcRebuildError(
|
| 1296 |
+
f"Group {group_index} {voice_id}: music line must end with a barline"
|
| 1297 |
+
)
|
| 1298 |
+
serialized_bars = [bar.strip() for bar in split_bars[:-1]]
|
| 1299 |
+
if any(not bar for bar in serialized_bars):
|
| 1300 |
+
raise AbcRebuildError(
|
| 1301 |
+
f"Group {group_index} {voice_id}: empty serialized measure"
|
| 1302 |
+
)
|
| 1303 |
+
bars = []
|
| 1304 |
+
for bar in serialized_bars:
|
| 1305 |
+
match = re.fullmatch(r"Z(?P<count>[1-4])?", bar)
|
| 1306 |
+
if match is None:
|
| 1307 |
+
bars.append(bar)
|
| 1308 |
+
continue
|
| 1309 |
+
if match.group("count") == "1":
|
| 1310 |
+
raise AbcRebuildError(
|
| 1311 |
+
f"Group {group_index} {voice_id}: Z1 must be written as Z"
|
| 1312 |
+
)
|
| 1313 |
+
bars.extend(["Z"] * int(match.group("count") or "1"))
|
| 1314 |
+
if not 1 <= len(bars) <= 4:
|
| 1315 |
+
raise AbcRebuildError(
|
| 1316 |
+
f"Group {group_index} {voice_id}: expected 1-4 semantic measures"
|
| 1317 |
+
)
|
| 1318 |
+
return cursor, fields, bars
|
| 1319 |
+
|
| 1320 |
+
|
| 1321 |
+
def validate_serialized_abc(text: str, score: RebuiltAbcScore) -> None:
|
| 1322 |
+
"""Validate invariants that must hold before any ABC is written."""
|
| 1323 |
+
|
| 1324 |
+
lines = text.splitlines()
|
| 1325 |
+
if not lines or lines[0] != "X:1":
|
| 1326 |
+
raise AbcRebuildError("ABC must start with X:1")
|
| 1327 |
+
if len(lines) < 2 or lines[1] != "T:":
|
| 1328 |
+
raise AbcRebuildError("ABC title must be fixed as empty T:")
|
| 1329 |
+
if any(line.startswith(("%abc-", "I:abc-creator")) for line in lines):
|
| 1330 |
+
raise AbcRebuildError("ABC must not contain version or creator metadata")
|
| 1331 |
+
if "%%MIDI gchordoff" in text:
|
| 1332 |
+
raise AbcRebuildError("ABC must not contain %%MIDI gchordoff")
|
| 1333 |
+
if "% ss2" in text:
|
| 1334 |
+
raise AbcRebuildError("ABC must not contain % ss2 metadata")
|
| 1335 |
+
header_voice_ids = [
|
| 1336 |
+
match.group(1)
|
| 1337 |
+
for line in lines
|
| 1338 |
+
if (match := re.match(r"^V: (Vocal|Ins) ", line))
|
| 1339 |
+
]
|
| 1340 |
+
if header_voice_ids != list(VOICE_IDS):
|
| 1341 |
+
raise AbcRebuildError(f"Expected fixed Vocal/Ins voice definitions, got {header_voice_ids!r}")
|
| 1342 |
+
|
| 1343 |
+
header_key_index = next(
|
| 1344 |
+
(
|
| 1345 |
+
index
|
| 1346 |
+
for index, line in enumerate(lines)
|
| 1347 |
+
if index > 0
|
| 1348 |
+
and line.startswith("K:")
|
| 1349 |
+
and any(
|
| 1350 |
+
header_index < index
|
| 1351 |
+
for header_index, header_line in enumerate(lines)
|
| 1352 |
+
if header_line.startswith("V: Ins ")
|
| 1353 |
+
)
|
| 1354 |
+
),
|
| 1355 |
+
None,
|
| 1356 |
+
)
|
| 1357 |
+
if header_key_index is None:
|
| 1358 |
+
raise AbcRebuildError("ABC header K: field is missing")
|
| 1359 |
+
first_measure = score.measures[0]
|
| 1360 |
+
expected_header_meter = (
|
| 1361 |
+
f"M:{first_measure.abc_numerator}/{first_measure.abc_denominator}"
|
| 1362 |
+
)
|
| 1363 |
+
header_meters = [
|
| 1364 |
+
line
|
| 1365 |
+
for line in lines[:header_key_index + 1]
|
| 1366 |
+
if line.startswith("M:")
|
| 1367 |
+
]
|
| 1368 |
+
if header_meters != [expected_header_meter]:
|
| 1369 |
+
raise AbcRebuildError(
|
| 1370 |
+
f"ABC header meters {header_meters!r} "
|
| 1371 |
+
f"!= {[expected_header_meter]!r}"
|
| 1372 |
+
)
|
| 1373 |
+
expected_header_key = f"K:{score.key_arr[first_measure.start_t]}"
|
| 1374 |
+
header_keys = [
|
| 1375 |
+
line
|
| 1376 |
+
for line in lines[:header_key_index + 1]
|
| 1377 |
+
if line.startswith("K:")
|
| 1378 |
+
]
|
| 1379 |
+
if header_keys != [expected_header_key]:
|
| 1380 |
+
raise AbcRebuildError(
|
| 1381 |
+
f"ABC header keys {header_keys!r} "
|
| 1382 |
+
f"!= {[expected_header_key]!r}"
|
| 1383 |
+
)
|
| 1384 |
+
|
| 1385 |
+
unit_denominator = abc_unit_denominator(score)
|
| 1386 |
+
expected_groups = _measure_groups(score)
|
| 1387 |
+
cursor = header_key_index + 1
|
| 1388 |
+
for group_index, group in enumerate(expected_groups):
|
| 1389 |
+
structure_labels = []
|
| 1390 |
+
while cursor < len(lines) and lines[cursor].startswith("% "):
|
| 1391 |
+
structure_labels.append(lines[cursor][2:].strip())
|
| 1392 |
+
cursor += 1
|
| 1393 |
+
if structure_labels != group.structure_labels:
|
| 1394 |
+
raise AbcRebuildError(
|
| 1395 |
+
f"Group {group_index}: structure labels "
|
| 1396 |
+
f"{structure_labels!r} != {group.structure_labels!r}"
|
| 1397 |
+
)
|
| 1398 |
+
|
| 1399 |
+
cursor, vocal_fields, vocal_bars = _parse_voice_group(
|
| 1400 |
+
lines,
|
| 1401 |
+
cursor,
|
| 1402 |
+
"Vocal",
|
| 1403 |
+
group_index,
|
| 1404 |
+
)
|
| 1405 |
+
cursor, ins_fields, ins_bars = _parse_voice_group(
|
| 1406 |
+
lines,
|
| 1407 |
+
cursor,
|
| 1408 |
+
"Ins",
|
| 1409 |
+
group_index,
|
| 1410 |
+
)
|
| 1411 |
+
if vocal_fields != ins_fields:
|
| 1412 |
+
raise AbcRebuildError(
|
| 1413 |
+
f"Group {group_index}: meter/key changes must be scoped to both voices"
|
| 1414 |
+
)
|
| 1415 |
+
first_measure = group.measures[0]
|
| 1416 |
+
expected_fields = {}
|
| 1417 |
+
if group.meter_changed:
|
| 1418 |
+
expected_fields["M"] = (
|
| 1419 |
+
f"{first_measure.abc_numerator}/{first_measure.abc_denominator}"
|
| 1420 |
+
)
|
| 1421 |
+
if group.key_changed:
|
| 1422 |
+
expected_fields["K"] = str(
|
| 1423 |
+
score.key_arr[first_measure.start_t]
|
| 1424 |
+
)
|
| 1425 |
+
if vocal_fields != expected_fields:
|
| 1426 |
+
raise AbcRebuildError(
|
| 1427 |
+
f"Group {group_index}: fields {vocal_fields!r} "
|
| 1428 |
+
f"!= required changes {expected_fields!r}"
|
| 1429 |
+
)
|
| 1430 |
+
if (
|
| 1431 |
+
len(vocal_bars) != len(group.measures)
|
| 1432 |
+
or len(ins_bars) != len(group.measures)
|
| 1433 |
+
):
|
| 1434 |
+
raise AbcRebuildError(
|
| 1435 |
+
f"Group {group_index}: both voices must contain "
|
| 1436 |
+
f"{len(group.measures)} measures"
|
| 1437 |
+
)
|
| 1438 |
+
|
| 1439 |
+
for bar_index, measure in enumerate(group.measures):
|
| 1440 |
+
expected_duration = (
|
| 1441 |
+
measure.abc_numerator
|
| 1442 |
+
* unit_denominator
|
| 1443 |
+
// measure.abc_denominator
|
| 1444 |
+
)
|
| 1445 |
+
for voice_id, bars in (
|
| 1446 |
+
("Vocal", vocal_bars),
|
| 1447 |
+
("Ins", ins_bars),
|
| 1448 |
+
):
|
| 1449 |
+
quoted_events, key_events = _parse_music_measure(
|
| 1450 |
+
bars[bar_index],
|
| 1451 |
+
expected_duration,
|
| 1452 |
+
f"measure {measure.index} {voice_id}",
|
| 1453 |
+
)
|
| 1454 |
+
expected_keys = _expected_measure_keys(
|
| 1455 |
+
score,
|
| 1456 |
+
measure,
|
| 1457 |
+
unit_denominator,
|
| 1458 |
+
)
|
| 1459 |
+
if key_events != expected_keys:
|
| 1460 |
+
raise AbcRebuildError(
|
| 1461 |
+
f"Measure {measure.index} {voice_id}: inline keys "
|
| 1462 |
+
f"{key_events!r} != {expected_keys!r}"
|
| 1463 |
+
)
|
| 1464 |
+
if voice_id == "Vocal":
|
| 1465 |
+
expected_chords = _expected_measure_chords(
|
| 1466 |
+
score,
|
| 1467 |
+
measure,
|
| 1468 |
+
unit_denominator,
|
| 1469 |
+
)
|
| 1470 |
+
if quoted_events != expected_chords:
|
| 1471 |
+
raise AbcRebuildError(
|
| 1472 |
+
f"Measure {measure.index}: chord symbols "
|
| 1473 |
+
f"{quoted_events!r} != {expected_chords!r}"
|
| 1474 |
+
)
|
| 1475 |
+
elif quoted_events:
|
| 1476 |
+
raise AbcRebuildError(
|
| 1477 |
+
f"Measure {measure.index}: chords must only be in Vocal"
|
| 1478 |
+
)
|
| 1479 |
+
if cursor != len(lines):
|
| 1480 |
+
raise AbcRebuildError(
|
| 1481 |
+
f"Unexpected trailing ABC body lines: {lines[cursor:cursor + 5]!r}"
|
| 1482 |
+
)
|
| 1483 |
+
|
| 1484 |
+
|
| 1485 |
+
def abc_paths_from_melody(
|
| 1486 |
+
melody_midi_path,
|
| 1487 |
+
output_path=None,
|
| 1488 |
+
*,
|
| 1489 |
+
melody_only=False,
|
| 1490 |
+
):
|
| 1491 |
+
melody = Path(melody_midi_path)
|
| 1492 |
+
if melody.name.endswith("_raw_full_melody.mid"):
|
| 1493 |
+
raise AbcRebuildError(f"Raw melody MIDI is not a valid ABC input: {melody}")
|
| 1494 |
+
if not melody.name.endswith("_melody.mid"):
|
| 1495 |
+
raise AbcRebuildError(f"Expected an exact *_melody.mid input, got {melody}")
|
| 1496 |
+
stem = melody.name[: -len("_melody.mid")]
|
| 1497 |
+
prefix = melody.with_name(stem)
|
| 1498 |
+
default_suffix = "_melody_only.abc" if melody_only else "_full.abc"
|
| 1499 |
+
return {
|
| 1500 |
+
"melody_midi": melody,
|
| 1501 |
+
"beats": Path(str(prefix) + "_beats.txt"),
|
| 1502 |
+
"chords": Path(str(prefix) + "_chords.txt"),
|
| 1503 |
+
"keys": Path(str(prefix) + "_keys.txt"),
|
| 1504 |
+
"structures": Path(str(prefix) + "_structures.txt"),
|
| 1505 |
+
"output": (
|
| 1506 |
+
Path(output_path)
|
| 1507 |
+
if output_path is not None
|
| 1508 |
+
else Path(str(prefix) + default_suffix)
|
| 1509 |
+
),
|
| 1510 |
+
}
|
| 1511 |
+
|
| 1512 |
+
|
| 1513 |
+
def preflight_exports(
|
| 1514 |
+
melody_midi_path,
|
| 1515 |
+
output_path=None,
|
| 1516 |
+
*,
|
| 1517 |
+
melody_only=False,
|
| 1518 |
+
):
|
| 1519 |
+
paths = abc_paths_from_melody(
|
| 1520 |
+
melody_midi_path,
|
| 1521 |
+
output_path=output_path,
|
| 1522 |
+
melody_only=melody_only,
|
| 1523 |
+
)
|
| 1524 |
+
required = {"melody_midi", "beats", "keys", "structures"}
|
| 1525 |
+
if not melody_only:
|
| 1526 |
+
required.add("chords")
|
| 1527 |
+
missing = [
|
| 1528 |
+
str(path)
|
| 1529 |
+
for name, path in paths.items()
|
| 1530 |
+
if name in required and not path.is_file()
|
| 1531 |
+
]
|
| 1532 |
+
if missing:
|
| 1533 |
+
raise FileNotFoundError(
|
| 1534 |
+
f"{paths['melody_midi']}: missing required companion file(s): {', '.join(missing)}"
|
| 1535 |
+
)
|
| 1536 |
+
return paths
|
| 1537 |
+
|
| 1538 |
+
|
| 1539 |
+
def generate_abc_from_exports(
|
| 1540 |
+
melody_midi_path,
|
| 1541 |
+
*,
|
| 1542 |
+
output_path=None,
|
| 1543 |
+
meter_conflict="infer",
|
| 1544 |
+
melody_only=False,
|
| 1545 |
+
):
|
| 1546 |
+
paths = preflight_exports(
|
| 1547 |
+
melody_midi_path,
|
| 1548 |
+
output_path=output_path,
|
| 1549 |
+
melody_only=melody_only,
|
| 1550 |
+
)
|
| 1551 |
+
score = build_rebuilt_abc_score(
|
| 1552 |
+
paths["melody_midi"],
|
| 1553 |
+
paths["beats"],
|
| 1554 |
+
paths["chords"],
|
| 1555 |
+
paths["keys"],
|
| 1556 |
+
paths["structures"],
|
| 1557 |
+
meter_conflict=meter_conflict,
|
| 1558 |
+
melody_only=melody_only,
|
| 1559 |
+
)
|
| 1560 |
+
return score_to_abc(score), score, paths
|
| 1561 |
+
|
| 1562 |
+
|
| 1563 |
+
def generate_abc_from_data(melody_midi, beats, chords, keys, structures, *,
|
| 1564 |
+
meter_conflict="infer", melody_only=False):
|
| 1565 |
+
"""Return validated ABC text and its score without filesystem access."""
|
| 1566 |
+
score = build_rebuilt_abc_score_from_data(
|
| 1567 |
+
melody_midi, beats, chords, keys, structures,
|
| 1568 |
+
meter_conflict=meter_conflict, melody_only=melody_only,
|
| 1569 |
+
)
|
| 1570 |
+
return score_to_abc(score), score
|
pipeline_sheetsage2.py
ADDED
|
@@ -0,0 +1,259 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Whole-song inference with cached overlap prefixes and right-hand audio context."""
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
import json
|
| 4 |
+
import re
|
| 5 |
+
import time
|
| 6 |
+
|
| 7 |
+
import torch
|
| 8 |
+
|
| 9 |
+
from .io_sheetsage2 import atomic_write_text
|
| 10 |
+
from .audio_sheetsage2 import SAMPLE_RATE, load_audio, slice_audio
|
| 11 |
+
from .generation_sheetsage2 import (
|
| 12 |
+
FULL_TASK_PROMPTS, build_overlap_prefix_tokens, constrained_prompt_generate,
|
| 13 |
+
decode_generated_tokens, event_time_map, stitched_window_events, write_window_tokens,
|
| 14 |
+
)
|
| 15 |
+
from .exports_sheetsage2 import export_result
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def _tensor_files(output_dir):
|
| 19 |
+
"""Inventory only files owned by the optional tensor exporter."""
|
| 20 |
+
files = set()
|
| 21 |
+
tensor_dir = output_dir / "tensors"
|
| 22 |
+
if tensor_dir.is_symlink():
|
| 23 |
+
return files
|
| 24 |
+
for directory in tensor_dir.glob("window-*"):
|
| 25 |
+
if re.fullmatch(r"window-\d{4,}", directory.name) and directory.is_dir() and not directory.is_symlink():
|
| 26 |
+
for path in directory.iterdir():
|
| 27 |
+
if path.is_file() and re.fullmatch(r"(?:index\.json|audio\.safetensors|(?:tokens|decoder)-\d{4,}\.safetensors)", path.name):
|
| 28 |
+
files.add(path)
|
| 29 |
+
return files
|
| 30 |
+
|
| 31 |
+
|
| 32 |
+
def sliding_window_plan(duration, window_seconds=300.0, overlap_seconds=200.0, lookahead_seconds=100.0):
|
| 33 |
+
if not duration > 0 or not window_seconds > 0:
|
| 34 |
+
raise ValueError("Duration and window length must be positive")
|
| 35 |
+
if not 0 <= lookahead_seconds <= overlap_seconds < window_seconds:
|
| 36 |
+
raise ValueError("Require 0 <= lookahead <= overlap < window length")
|
| 37 |
+
hop = window_seconds - overlap_seconds
|
| 38 |
+
start, accepted = 0.0, 0.0
|
| 39 |
+
result = []
|
| 40 |
+
while True:
|
| 41 |
+
last = start + window_seconds >= duration - 1e-6
|
| 42 |
+
accept_end = duration if last else start + window_seconds - lookahead_seconds
|
| 43 |
+
result.append(dict(start=start, end=min(duration, start + window_seconds),
|
| 44 |
+
accept_start=accepted, accept_end=accept_end, prefix_end=accepted,
|
| 45 |
+
generation_stop=None if last else window_seconds - lookahead_seconds))
|
| 46 |
+
if last:
|
| 47 |
+
return result
|
| 48 |
+
accepted = accept_end
|
| 49 |
+
start = min(start + hop, duration - window_seconds)
|
| 50 |
+
|
| 51 |
+
|
| 52 |
+
class Transcriber:
|
| 53 |
+
def __init__(self, model, dtype="bf16"):
|
| 54 |
+
self.model = model
|
| 55 |
+
self.device = next(model.parameters()).device
|
| 56 |
+
self.dtype = {"bf16": torch.bfloat16, "fp32": None}[dtype]
|
| 57 |
+
|
| 58 |
+
@torch.inference_mode()
|
| 59 |
+
def analyze(self, audio_path, output_dir=None, *, prompts=FULL_TASK_PROMPTS,
|
| 60 |
+
sampling_rate=None, max_seconds=None, preset="default",
|
| 61 |
+
overlap_seconds=None, lookahead_seconds=None, progress=None,
|
| 62 |
+
export_logits=False, export_scores=False, export_embeddings=False,
|
| 63 |
+
output_hidden_states=False, melody_only=False):
|
| 64 |
+
def report(stage, **fields):
|
| 65 |
+
if progress:
|
| 66 |
+
progress(dict(stage=stage, **fields))
|
| 67 |
+
|
| 68 |
+
if not isinstance(melody_only, bool):
|
| 69 |
+
raise ValueError("melody_only must be True or False")
|
| 70 |
+
if preset not in ("default", "paper"):
|
| 71 |
+
raise ValueError("preset must be default or paper")
|
| 72 |
+
defaults = (200.0, 100.0) if preset == "default" else (100.0, 0.0)
|
| 73 |
+
overlap_seconds = defaults[0] if overlap_seconds is None else overlap_seconds
|
| 74 |
+
lookahead_seconds = defaults[1] if lookahead_seconds is None else lookahead_seconds
|
| 75 |
+
if preset == "paper" and (overlap_seconds, lookahead_seconds) != defaults:
|
| 76 |
+
raise ValueError("paper preset fixes overlap=100 and lookahead=0")
|
| 77 |
+
started = time.monotonic()
|
| 78 |
+
output_dir = Path(output_dir) if output_dir is not None else None
|
| 79 |
+
if output_dir is not None:
|
| 80 |
+
output_dir.mkdir(parents=True, exist_ok=True)
|
| 81 |
+
previous_tensors = _tensor_files(output_dir) if output_dir is not None else set()
|
| 82 |
+
current_tensors = set()
|
| 83 |
+
tensor_results = []
|
| 84 |
+
prompts = self.model.tokenizer.normalize_prompts(prompts)
|
| 85 |
+
if "timestamp" not in prompts:
|
| 86 |
+
raise ValueError("timestamp is required to export timed annotations")
|
| 87 |
+
report("audio")
|
| 88 |
+
audio = load_audio(audio_path, sampling_rate=sampling_rate, max_seconds=max_seconds, preset=preset)
|
| 89 |
+
duration = len(audio) / SAMPLE_RATE
|
| 90 |
+
window_length = float(self.model.hparams.input_audio_length)
|
| 91 |
+
plan = sliding_window_plan(duration, window_length, overlap_seconds, lookahead_seconds)
|
| 92 |
+
if self.device.type == "cuda":
|
| 93 |
+
torch.cuda.reset_peak_memory_stats(self.device)
|
| 94 |
+
stitched, records, warnings = [], [], []
|
| 95 |
+
for index, window in enumerate(plan):
|
| 96 |
+
report("encoding", window=index + 1, windows=len(plan), start=window["start"])
|
| 97 |
+
segment = slice_audio(audio, window["start"], window_length)[None].to(self.device)
|
| 98 |
+
prefix, base = None, 0
|
| 99 |
+
if index:
|
| 100 |
+
_, prefix, base = build_overlap_prefix_tokens(
|
| 101 |
+
stitched, self.model.tokenizer, prompts, window["start"], window["prefix_end"])
|
| 102 |
+
if prefix is not None:
|
| 103 |
+
if preset == "paper" and len(prefix) >= self.model.max_output_seq_len - 16:
|
| 104 |
+
warnings.append(f"Window {index + 1}: overlap prefix exceeded the context")
|
| 105 |
+
prefix = None
|
| 106 |
+
elif preset != "paper" and len(prefix) >= self.model.max_output_seq_len - 128:
|
| 107 |
+
raise ValueError("Overlap prefix fills the context; reduce overlap_seconds")
|
| 108 |
+
tick = time.monotonic()
|
| 109 |
+
memory = None
|
| 110 |
+
tensor_record = None
|
| 111 |
+
if export_embeddings or output_hidden_states or export_logits or export_scores:
|
| 112 |
+
from .tensors_sheetsage2 import WindowTensorWriter
|
| 113 |
+
tensor_record = WindowTensorWriter(output_dir, index, window, export_logits, export_scores)
|
| 114 |
+
if export_embeddings or output_hidden_states:
|
| 115 |
+
memory = tensor_record.audio_features(self.model, segment, self.dtype, output_hidden_states)
|
| 116 |
+
tokens = constrained_prompt_generate(
|
| 117 |
+
self.model, segment, prompts, self.model.max_output_seq_len,
|
| 118 |
+
prefix_tokens=prefix, autocast_dtype=self.dtype,
|
| 119 |
+
stop_time_seconds=(None if preset == "paper" else
|
| 120 |
+
window["generation_stop"] if window["generation_stop"] is not None else
|
| 121 |
+
min(duration - window["start"], window_length)),
|
| 122 |
+
memory=memory,
|
| 123 |
+
step_callback=tensor_record.capture if tensor_record is not None and (export_logits or export_scores) else None,
|
| 124 |
+
progress_callback=lambda n: report("decoding", window=index + 1, windows=len(plan), tokens=n),
|
| 125 |
+
)
|
| 126 |
+
if tensor_record is not None:
|
| 127 |
+
if export_embeddings:
|
| 128 |
+
tensor_record.decoder_features(self.model, segment, tokens, self.dtype, memory)
|
| 129 |
+
tensor_data = tensor_record.finish()
|
| 130 |
+
if tensor_data is not None:
|
| 131 |
+
tensor_results.append(tensor_data)
|
| 132 |
+
if tensor_record.directory is not None:
|
| 133 |
+
current_tensors.update(tensor_record.directory / name for name in tensor_record.files)
|
| 134 |
+
current_tensors.add(tensor_record.directory / "index.json")
|
| 135 |
+
if len(tokens) > self.model.max_output_seq_len:
|
| 136 |
+
warnings.append(f"Window {index + 1} reached the token limit; inspect its token coverage")
|
| 137 |
+
decoded, warning = decode_generated_tokens(self.model.tokenizer, tokens, Path(audio_path).stem if isinstance(audio_path, (str, Path)) else "audio", index)
|
| 138 |
+
if warning:
|
| 139 |
+
warnings.append(warning["error"])
|
| 140 |
+
lookup = event_time_map(decoded, window_length)
|
| 141 |
+
accepted = stitched_window_events(
|
| 142 |
+
decoded, lookup, window["start"], window["accept_start"], window["accept_end"],
|
| 143 |
+
duration, index, global_subbeat_base=base or 0,
|
| 144 |
+
)
|
| 145 |
+
stitched.extend(accepted)
|
| 146 |
+
if preset == "paper":
|
| 147 |
+
stitched.sort(key=lambda e: (float(e.get("time", 0)), int(e.get("window_index", 0)),
|
| 148 |
+
int(e.get("source_subbeat", e["subbeat"]))))
|
| 149 |
+
record = dict(window, window_index=index, prefix_tokens=0 if prefix is None else len(prefix),
|
| 150 |
+
tokens=tokens, events=len(decoded["events"]), accepted_events=len(accepted),
|
| 151 |
+
elapsed_seconds=time.monotonic() - tick)
|
| 152 |
+
records.append(record)
|
| 153 |
+
report("window_complete", window=index + 1, windows=len(plan), tokens=len(tokens))
|
| 154 |
+
if preset != "paper":
|
| 155 |
+
stitched.sort(key=lambda e: (e["time"], e["global_subbeat"]))
|
| 156 |
+
decoded = dict(schema_version=self.model.tokenizer.schema_version,
|
| 157 |
+
prompts=list(prompts), events=stitched, has_eos=True)
|
| 158 |
+
report("notation")
|
| 159 |
+
exported = export_result(decoded, self.model.tokenizer, output_dir, duration,
|
| 160 |
+
paper=preset == "paper", melody_only=melody_only)
|
| 161 |
+
payload = exported.pop("payload")
|
| 162 |
+
if output_dir is not None:
|
| 163 |
+
write_window_tokens(output_dir / "tokens.txt", records, self.model.tokenizer)
|
| 164 |
+
atomic_write_text(output_dir / "tokens.json", json.dumps([
|
| 165 |
+
dict(r, tokens=r["tokens"].tolist()) for r in records
|
| 166 |
+
]))
|
| 167 |
+
result = dict(
|
| 168 |
+
audio=Path(audio_path).name if isinstance(audio_path, (str, Path)) else "audio",
|
| 169 |
+
duration_seconds=duration, prompts=list(prompts), preset=preset, melody_only=melody_only,
|
| 170 |
+
dtype="bf16" if self.dtype and self.device.type == "cuda" else "fp32",
|
| 171 |
+
window_seconds=window_length, overlap_seconds=overlap_seconds, lookahead_seconds=lookahead_seconds,
|
| 172 |
+
windows=[dict(r, tokens=len(r["tokens"])) for r in records],
|
| 173 |
+
elapsed_seconds=time.monotonic() - started, warnings=warnings,
|
| 174 |
+
peak_gpu_mib=torch.cuda.max_memory_allocated(self.device) / 1024**2 if self.device.type == "cuda" else 0,
|
| 175 |
+
**exported,
|
| 176 |
+
)
|
| 177 |
+
if output_dir is not None:
|
| 178 |
+
atomic_write_text(output_dir / "result.json", json.dumps(result, indent=2))
|
| 179 |
+
for path in previous_tensors - current_tensors:
|
| 180 |
+
path.unlink(missing_ok=True)
|
| 181 |
+
for directory in {path.parent for path in previous_tensors - current_tensors}:
|
| 182 |
+
if directory.is_dir() and not any(directory.iterdir()):
|
| 183 |
+
directory.rmdir()
|
| 184 |
+
report("complete", **exported)
|
| 185 |
+
return dict(result, num_events=result["events"], **payload,
|
| 186 |
+
tokens=[r["tokens"] for r in records], tensors=tensor_results,
|
| 187 |
+
_metadata=result)
|
| 188 |
+
|
| 189 |
+
|
| 190 |
+
@torch.inference_mode()
|
| 191 |
+
def transcribe(model, audio, output_dir=None, *, render_audio=False, render_score=False,
|
| 192 |
+
render_parts=("mix",), dtype="bf16", melody_only=False, **kwargs):
|
| 193 |
+
"""Transcribe a path, encoded audio bytes, binary stream, or waveform.
|
| 194 |
+
|
| 195 |
+
Arrays use channels-first layout and require sampling_rate. With
|
| 196 |
+
output_dir=None, all audio/results stay in memory: ABC is text, MIDI is
|
| 197 |
+
bytes, events are a list, and optional features are CPU tensors grouped by
|
| 198 |
+
window. A directory saves the standard output files and writes optional
|
| 199 |
+
tensors to disk instead of keeping them in memory. Optional
|
| 200 |
+
rendering also returns bytes/text when no output directory is given.
|
| 201 |
+
melody_only=True omits chords from ABC and MIDI playback, retaining both
|
| 202 |
+
melody voices and the original decoded events/LAB annotations. If the
|
| 203 |
+
requested melody-only ABC cannot be built, raises RuntimeError with the
|
| 204 |
+
completed transcription available as error.result.
|
| 205 |
+
"""
|
| 206 |
+
if not isinstance(melody_only, bool):
|
| 207 |
+
raise ValueError("melody_only must be True or False")
|
| 208 |
+
output_dir = Path(output_dir) if output_dir is not None else None
|
| 209 |
+
if render_audio or render_score:
|
| 210 |
+
from .rendering_sheetsage2 import render_outputs, render_memory, validate_render_options
|
| 211 |
+
formats, render_parts = validate_render_options(audio=render_audio, score=render_score, parts=render_parts)
|
| 212 |
+
result = Transcriber(model, dtype=dtype).analyze(audio, output_dir, melody_only=melody_only, **kwargs)
|
| 213 |
+
metadata = result.pop("_metadata", None)
|
| 214 |
+
if metadata is None:
|
| 215 |
+
metadata = dict(result)
|
| 216 |
+
if render_audio or render_score:
|
| 217 |
+
try:
|
| 218 |
+
assets = Path(__file__).with_name("render_assets")
|
| 219 |
+
if not assets.is_dir():
|
| 220 |
+
source = getattr(model, "_source_snapshot", Path(model.config._name_or_path))
|
| 221 |
+
if (source / "render_assets").is_dir():
|
| 222 |
+
assets = source / "render_assets"
|
| 223 |
+
else:
|
| 224 |
+
from huggingface_hub import snapshot_download
|
| 225 |
+
assets = Path(snapshot_download(model.config._name_or_path,
|
| 226 |
+
revision=model.config._commit_hash, allow_patterns=["render_assets/**"],
|
| 227 |
+
**getattr(model, "_hub_resource_options", {}))) / "render_assets"
|
| 228 |
+
# A missing score does not prevent a requested MIDI piano preview.
|
| 229 |
+
rendered = None
|
| 230 |
+
if render_audio or not result.get("abc_error"):
|
| 231 |
+
options = dict(audio=render_audio, score=() if result.get("abc_error") else formats,
|
| 232 |
+
parts=render_parts, assets_dir=assets)
|
| 233 |
+
if output_dir is None:
|
| 234 |
+
rendered = render_memory(midi=result["midi"], abc=result["abc"],
|
| 235 |
+
duration=result["duration_seconds"], **options)
|
| 236 |
+
else:
|
| 237 |
+
rendered = render_outputs(input_dir=output_dir, output_dir=output_dir, **options)
|
| 238 |
+
result["rendered"] = rendered
|
| 239 |
+
if output_dir is not None:
|
| 240 |
+
metadata["rendered"] = rendered
|
| 241 |
+
if formats and result.get("abc_error"):
|
| 242 |
+
raise ValueError(f"ABC unavailable: {result['abc_error']}")
|
| 243 |
+
except Exception as exc:
|
| 244 |
+
result["render_error"] = str(exc)
|
| 245 |
+
metadata["render_error"] = str(exc)
|
| 246 |
+
if output_dir is not None:
|
| 247 |
+
atomic_write_text(output_dir / "result.json", json.dumps(metadata, indent=2))
|
| 248 |
+
location = f"saved to {output_dir.resolve()}" if output_dir is not None else "completed in memory"
|
| 249 |
+
error = RuntimeError(f"Transcription {location}; rendering failed: {exc}")
|
| 250 |
+
error.result = result
|
| 251 |
+
raise error from exc
|
| 252 |
+
if output_dir is not None:
|
| 253 |
+
atomic_write_text(output_dir / "result.json", json.dumps(metadata, indent=2))
|
| 254 |
+
if melody_only and not result.get("abc"):
|
| 255 |
+
reason = result.get("abc_error") or "No ABC score was produced"
|
| 256 |
+
error = RuntimeError(f"Melody-only ABC unavailable: {reason}")
|
| 257 |
+
error.result = result
|
| 258 |
+
raise error
|
| 259 |
+
return result
|
processing_sheetsage2.py
ADDED
|
@@ -0,0 +1,109 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Waveform preparation and symbolic decoding for SheetSage2."""
|
| 2 |
+
|
| 3 |
+
import numbers
|
| 4 |
+
|
| 5 |
+
import torch
|
| 6 |
+
import torchaudio
|
| 7 |
+
from transformers import ProcessorMixin
|
| 8 |
+
from transformers.feature_extraction_utils import BatchFeature
|
| 9 |
+
|
| 10 |
+
from .tokenization_sheetsage2 import SheetSage2Tokenizer
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def _hf_relative_dependencies():
|
| 14 |
+
# Keep standalone AutoProcessor loading compatible with Transformers 4.45.
|
| 15 |
+
from .durations_sheetsage2 import DURATION_TEMPLATES
|
| 16 |
+
from .labels_sheetsage2 import STRUCTURE_LABELS
|
| 17 |
+
from .schema_sheetsage2 import get_prompt_multitask_schema
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
class SheetSage2Processor(ProcessorMixin):
|
| 21 |
+
"""Prepare mono waveforms without amplitude normalization.
|
| 22 |
+
|
| 23 |
+
Pass one mono waveform, a batch tensor, or a list of mono waveforms. Use
|
| 24 |
+
``model.transcribe`` for files and audio longer than one model window.
|
| 25 |
+
"""
|
| 26 |
+
|
| 27 |
+
attributes = []
|
| 28 |
+
valid_kwargs = ["sampling_rate", "window_seconds", "time_hz", "schema_version", "tokenizer_fingerprint"]
|
| 29 |
+
|
| 30 |
+
def __init__(self, sampling_rate=24000, window_seconds=300.0, time_hz=100,
|
| 31 |
+
schema_version="v1", tokenizer_fingerprint="5ba3325af0344c7f", **kwargs):
|
| 32 |
+
super().__init__(**{key: value for key, value in kwargs.items() if key == "chat_template"})
|
| 33 |
+
self.sampling_rate = int(sampling_rate)
|
| 34 |
+
self.window_seconds = float(window_seconds)
|
| 35 |
+
self.time_hz = int(time_hz)
|
| 36 |
+
self.schema_version = str(schema_version)
|
| 37 |
+
self.tokenizer_fingerprint = str(tokenizer_fingerprint)
|
| 38 |
+
self.tokenizer = SheetSage2Tokenizer(
|
| 39 |
+
self.window_seconds, self.time_hz, self.schema_version,
|
| 40 |
+
expected_fingerprint=self.tokenizer_fingerprint,
|
| 41 |
+
)
|
| 42 |
+
|
| 43 |
+
@property
|
| 44 |
+
def model_input_names(self):
|
| 45 |
+
return ["input_values", "attention_mask"]
|
| 46 |
+
|
| 47 |
+
@classmethod
|
| 48 |
+
def from_model_config(cls, config):
|
| 49 |
+
return cls(config.sampling_rate, config.input_audio_length, config.time_hz,
|
| 50 |
+
config.tokenizer_schema_version, config.tokenizer_fingerprint)
|
| 51 |
+
|
| 52 |
+
def __call__(self, audio, sampling_rate=None, padding=True, return_tensors="pt",
|
| 53 |
+
return_attention_mask=True):
|
| 54 |
+
source_rate = self.sampling_rate if sampling_rate is None else int(sampling_rate)
|
| 55 |
+
if source_rate <= 0:
|
| 56 |
+
raise ValueError("sampling_rate must be positive.")
|
| 57 |
+
if isinstance(audio, (str, bytes)):
|
| 58 |
+
raise ValueError("Pass waveform samples here; use model.transcribe for audio files.")
|
| 59 |
+
if isinstance(audio, (list, tuple)) and audio and not isinstance(audio[0], numbers.Number):
|
| 60 |
+
waveforms = [torch.as_tensor(value) for value in audio]
|
| 61 |
+
else:
|
| 62 |
+
value = torch.as_tensor(audio)
|
| 63 |
+
waveforms = list(value) if value.ndim == 2 else [value]
|
| 64 |
+
if not waveforms:
|
| 65 |
+
raise ValueError("Provide at least one waveform.")
|
| 66 |
+
prepared = []
|
| 67 |
+
maximum = round(self.window_seconds * self.sampling_rate)
|
| 68 |
+
for waveform in waveforms:
|
| 69 |
+
if waveform.ndim != 1 or not waveform.is_floating_point() or not torch.isfinite(waveform).all():
|
| 70 |
+
raise ValueError("Each waveform must be a one-dimensional finite floating-point array.")
|
| 71 |
+
waveform = waveform.to(dtype=torch.float32)
|
| 72 |
+
if source_rate != self.sampling_rate:
|
| 73 |
+
waveform = torchaudio.functional.resample(waveform, source_rate, self.sampling_rate)
|
| 74 |
+
if waveform.numel() < 1025:
|
| 75 |
+
raise ValueError("Each waveform must contain at least 1025 samples at 24 kHz.")
|
| 76 |
+
if waveform.numel() > maximum:
|
| 77 |
+
raise ValueError("Audio exceeds one model window; use model.transcribe for whole songs.")
|
| 78 |
+
prepared.append(waveform)
|
| 79 |
+
if len({str(value.device) for value in prepared}) != 1:
|
| 80 |
+
raise ValueError("All waveforms in a batch must be on the same device.")
|
| 81 |
+
if padding is False:
|
| 82 |
+
length = max(value.numel() for value in prepared)
|
| 83 |
+
if any(value.numel() != length for value in prepared):
|
| 84 |
+
raise ValueError("Use padding=True for waveforms of different lengths.")
|
| 85 |
+
elif padding is True or padding == "max_length":
|
| 86 |
+
length = maximum
|
| 87 |
+
elif padding == "longest":
|
| 88 |
+
length = max(value.numel() for value in prepared)
|
| 89 |
+
else:
|
| 90 |
+
raise ValueError("padding must be True, False, 'max_length', or 'longest'.")
|
| 91 |
+
lengths = torch.tensor([value.numel() for value in prepared], device=prepared[0].device)
|
| 92 |
+
values = torch.stack([torch.nn.functional.pad(value, (0, length - value.numel())) for value in prepared])
|
| 93 |
+
data = {"input_values": values}
|
| 94 |
+
if return_attention_mask:
|
| 95 |
+
data["attention_mask"] = (torch.arange(length, device=values.device)[None] < lengths[:, None]).long()
|
| 96 |
+
if return_tensors == "np":
|
| 97 |
+
data = {key: value.cpu().numpy() for key, value in data.items()}
|
| 98 |
+
elif return_tensors not in {None, "pt"}:
|
| 99 |
+
raise ValueError("return_tensors must be 'pt', 'np', or None.")
|
| 100 |
+
return BatchFeature(data=data)
|
| 101 |
+
|
| 102 |
+
def decode(self, token_ids, strict=True):
|
| 103 |
+
return self.tokenizer.decode_sequence(token_ids, strict=strict)
|
| 104 |
+
|
| 105 |
+
def batch_decode(self, sequences, strict=True):
|
| 106 |
+
return [self.decode(tokens, strict=strict) for tokens in sequences]
|
| 107 |
+
|
| 108 |
+
|
| 109 |
+
SheetSage2Processor.register_for_auto_class()
|
processor_config.json
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"processor_class": "SheetSage2Processor",
|
| 3 |
+
"sampling_rate": 24000,
|
| 4 |
+
"window_seconds": 300.0,
|
| 5 |
+
"time_hz": 100,
|
| 6 |
+
"schema_version": "v1",
|
| 7 |
+
"tokenizer_fingerprint": "5ba3325af0344c7f",
|
| 8 |
+
"auto_map": {
|
| 9 |
+
"AutoProcessor": "processing_sheetsage2.SheetSage2Processor"
|
| 10 |
+
}
|
| 11 |
+
}
|
render.py
ADDED
|
@@ -0,0 +1,45 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Render existing SheetSage2 MIDI and ABC outputs without loading a model."""
|
| 2 |
+
import argparse
|
| 3 |
+
import json
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
import sys
|
| 6 |
+
|
| 7 |
+
SCRIPT_DIR = Path(__file__).absolute().parent
|
| 8 |
+
sys.path.insert(0, str(SCRIPT_DIR))
|
| 9 |
+
|
| 10 |
+
from rendering_sheetsage2 import render_outputs
|
| 11 |
+
|
| 12 |
+
|
| 13 |
+
def main():
|
| 14 |
+
parser = argparse.ArgumentParser(description=__doc__)
|
| 15 |
+
parser.add_argument("--input", help="Transcription output directory.")
|
| 16 |
+
parser.add_argument("--midi", help="MIDI file with original note timing.")
|
| 17 |
+
parser.add_argument("--abc", help="ABC score file.")
|
| 18 |
+
parser.add_argument("--output", help="Output directory; defaults to --input or rendered/.")
|
| 19 |
+
parser.add_argument("--audio", action="store_true", help="Render piano WAV.")
|
| 20 |
+
parser.add_argument("--score", default="", help="Comma-separated pdf,svg,png.")
|
| 21 |
+
parser.add_argument("--parts", default="mix", help="Comma-separated mix,melody,vocal,instrumental,chords, or all (audio only).")
|
| 22 |
+
parser.add_argument("--duration", type=float, help="Minimum audio duration in seconds; preserves trailing silence.")
|
| 23 |
+
args = parser.parse_args()
|
| 24 |
+
audio = args.audio
|
| 25 |
+
score = args.score
|
| 26 |
+
if not audio and not score:
|
| 27 |
+
if args.input:
|
| 28 |
+
audio, score = True, "pdf"
|
| 29 |
+
else:
|
| 30 |
+
audio, score = bool(args.midi), "pdf" if args.abc else ""
|
| 31 |
+
if not args.input and not args.midi and not args.abc:
|
| 32 |
+
parser.error("Provide --input, --midi, or --abc.")
|
| 33 |
+
try:
|
| 34 |
+
result = render_outputs(args.input, midi=args.midi, abc=args.abc,
|
| 35 |
+
output_dir=args.output, audio=audio, score=score,
|
| 36 |
+
parts=args.parts, duration=args.duration)
|
| 37 |
+
except (ValueError, OSError, RuntimeError) as exc:
|
| 38 |
+
print(f"Rendering failed: {exc}", file=sys.stderr)
|
| 39 |
+
return 1
|
| 40 |
+
print(json.dumps(result, indent=2))
|
| 41 |
+
return 0
|
| 42 |
+
|
| 43 |
+
|
| 44 |
+
if __name__ == "__main__":
|
| 45 |
+
raise SystemExit(main())
|
render_assets/DejaVuSans.ttf
ADDED
|
@@ -0,0 +1,3 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
version https://git-lfs.github.com/spec/v1
|
| 2 |
+
oid sha256:7da195a74c55bef988d0d48f9508bd5d849425c1770dba5d7bfc6ce9ed848954
|
| 3 |
+
size 757076
|
render_assets/LICENSE.abcjs
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Copyright (c) 2009-2024 Paul Rosen and Gregory Dyke
|
| 2 |
+
|
| 3 |
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 4 |
+
of this software and associated documentation files (the "Software"), to deal
|
| 5 |
+
in the Software without restriction, including without limitation the rights
|
| 6 |
+
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
| 7 |
+
copies of the Software, and to permit persons to whom the Software is
|
| 8 |
+
furnished to do so, subject to the following conditions:
|
| 9 |
+
|
| 10 |
+
The above copyright notice and this permission notice shall be included in
|
| 11 |
+
all copies or substantial portions of the Software.
|
| 12 |
+
|
| 13 |
+
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
|
| 14 |
+
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
|
| 15 |
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
|
| 16 |
+
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
|
| 17 |
+
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
|
| 18 |
+
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN
|
| 19 |
+
THE SOFTWARE.
|
| 20 |
+
|
| 21 |
+
**This text is from: http://opensource.org/licenses/MIT**
|
render_assets/LICENSE.font
ADDED
|
@@ -0,0 +1,187 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
Fonts are (c) Bitstream (see below). DejaVu changes are in public domain.
|
| 2 |
+
Glyphs imported from Arev fonts are (c) Tavmjong Bah (see below)
|
| 3 |
+
|
| 4 |
+
|
| 5 |
+
Bitstream Vera Fonts Copyright
|
| 6 |
+
------------------------------
|
| 7 |
+
|
| 8 |
+
Copyright (c) 2003 by Bitstream, Inc. All Rights Reserved. Bitstream Vera is
|
| 9 |
+
a trademark of Bitstream, Inc.
|
| 10 |
+
|
| 11 |
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 12 |
+
of the fonts accompanying this license ("Fonts") and associated
|
| 13 |
+
documentation files (the "Font Software"), to reproduce and distribute the
|
| 14 |
+
Font Software, including without limitation the rights to use, copy, merge,
|
| 15 |
+
publish, distribute, and/or sell copies of the Font Software, and to permit
|
| 16 |
+
persons to whom the Font Software is furnished to do so, subject to the
|
| 17 |
+
following conditions:
|
| 18 |
+
|
| 19 |
+
The above copyright and trademark notices and this permission notice shall
|
| 20 |
+
be included in all copies of one or more of the Font Software typefaces.
|
| 21 |
+
|
| 22 |
+
The Font Software may be modified, altered, or added to, and in particular
|
| 23 |
+
the designs of glyphs or characters in the Fonts may be modified and
|
| 24 |
+
additional glyphs or characters may be added to the Fonts, only if the fonts
|
| 25 |
+
are renamed to names not containing either the words "Bitstream" or the word
|
| 26 |
+
"Vera".
|
| 27 |
+
|
| 28 |
+
This License becomes null and void to the extent applicable to Fonts or Font
|
| 29 |
+
Software that has been modified and is distributed under the "Bitstream
|
| 30 |
+
Vera" names.
|
| 31 |
+
|
| 32 |
+
The Font Software may be sold as part of a larger software package but no
|
| 33 |
+
copy of one or more of the Font Software typefaces may be sold by itself.
|
| 34 |
+
|
| 35 |
+
THE FONT SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS
|
| 36 |
+
OR IMPLIED, INCLUDING BUT NOT LIMITED TO ANY WARRANTIES OF MERCHANTABILITY,
|
| 37 |
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT OF COPYRIGHT, PATENT,
|
| 38 |
+
TRADEMARK, OR OTHER RIGHT. IN NO EVENT SHALL BITSTREAM OR THE GNOME
|
| 39 |
+
FOUNDATION BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, INCLUDING
|
| 40 |
+
ANY GENERAL, SPECIAL, INDIRECT, INCIDENTAL, OR CONSEQUENTIAL DAMAGES,
|
| 41 |
+
WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF
|
| 42 |
+
THE USE OR INABILITY TO USE THE FONT SOFTWARE OR FROM OTHER DEALINGS IN THE
|
| 43 |
+
FONT SOFTWARE.
|
| 44 |
+
|
| 45 |
+
Except as contained in this notice, the names of Gnome, the Gnome
|
| 46 |
+
Foundation, and Bitstream Inc., shall not be used in advertising or
|
| 47 |
+
otherwise to promote the sale, use or other dealings in this Font Software
|
| 48 |
+
without prior written authorization from the Gnome Foundation or Bitstream
|
| 49 |
+
Inc., respectively. For further information, contact: fonts at gnome dot
|
| 50 |
+
org.
|
| 51 |
+
|
| 52 |
+
Arev Fonts Copyright
|
| 53 |
+
------------------------------
|
| 54 |
+
|
| 55 |
+
Copyright (c) 2006 by Tavmjong Bah. All Rights Reserved.
|
| 56 |
+
|
| 57 |
+
Permission is hereby granted, free of charge, to any person obtaining
|
| 58 |
+
a copy of the fonts accompanying this license ("Fonts") and
|
| 59 |
+
associated documentation files (the "Font Software"), to reproduce
|
| 60 |
+
and distribute the modifications to the Bitstream Vera Font Software,
|
| 61 |
+
including without limitation the rights to use, copy, merge, publish,
|
| 62 |
+
distribute, and/or sell copies of the Font Software, and to permit
|
| 63 |
+
persons to whom the Font Software is furnished to do so, subject to
|
| 64 |
+
the following conditions:
|
| 65 |
+
|
| 66 |
+
The above copyright and trademark notices and this permission notice
|
| 67 |
+
shall be included in all copies of one or more of the Font Software
|
| 68 |
+
typefaces.
|
| 69 |
+
|
| 70 |
+
The Font Software may be modified, altered, or added to, and in
|
| 71 |
+
particular the designs of glyphs or characters in the Fonts may be
|
| 72 |
+
modified and additional glyphs or characters may be added to the
|
| 73 |
+
Fonts, only if the fonts are renamed to names not containing either
|
| 74 |
+
the words "Tavmjong Bah" or the word "Arev".
|
| 75 |
+
|
| 76 |
+
This License becomes null and void to the extent applicable to Fonts
|
| 77 |
+
or Font Software that has been modified and is distributed under the
|
| 78 |
+
"Tavmjong Bah Arev" names.
|
| 79 |
+
|
| 80 |
+
The Font Software may be sold as part of a larger software package but
|
| 81 |
+
no copy of one or more of the Font Software typefaces may be sold by
|
| 82 |
+
itself.
|
| 83 |
+
|
| 84 |
+
THE FONT SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND,
|
| 85 |
+
EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO ANY WARRANTIES OF
|
| 86 |
+
MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT
|
| 87 |
+
OF COPYRIGHT, PATENT, TRADEMARK, OR OTHER RIGHT. IN NO EVENT SHALL
|
| 88 |
+
TAVMJONG BAH BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY,
|
| 89 |
+
INCLUDING ANY GENERAL, SPECIAL, INDIRECT, INCIDENTAL, OR CONSEQUENTIAL
|
| 90 |
+
DAMAGES, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING
|
| 91 |
+
FROM, OUT OF THE USE OR INABILITY TO USE THE FONT SOFTWARE OR FROM
|
| 92 |
+
OTHER DEALINGS IN THE FONT SOFTWARE.
|
| 93 |
+
|
| 94 |
+
Except as contained in this notice, the name of Tavmjong Bah shall not
|
| 95 |
+
be used in advertising or otherwise to promote the sale, use or other
|
| 96 |
+
dealings in this Font Software without prior written authorization
|
| 97 |
+
from Tavmjong Bah. For further information, contact: tavmjong @ free
|
| 98 |
+
. fr.
|
| 99 |
+
|
| 100 |
+
TeX Gyre DJV Math
|
| 101 |
+
-----------------
|
| 102 |
+
Fonts are (c) Bitstream (see below). DejaVu changes are in public domain.
|
| 103 |
+
|
| 104 |
+
Math extensions done by B. Jackowski, P. Strzelczyk and P. Pianowski
|
| 105 |
+
(on behalf of TeX users groups) are in public domain.
|
| 106 |
+
|
| 107 |
+
Letters imported from Euler Fraktur from AMSfonts are (c) American
|
| 108 |
+
Mathematical Society (see below).
|
| 109 |
+
Bitstream Vera Fonts Copyright
|
| 110 |
+
Copyright (c) 2003 by Bitstream, Inc. All Rights Reserved. Bitstream Vera
|
| 111 |
+
is a trademark of Bitstream, Inc.
|
| 112 |
+
|
| 113 |
+
Permission is hereby granted, free of charge, to any person obtaining a copy
|
| 114 |
+
of the fonts accompanying this license (“Fonts”) and associated
|
| 115 |
+
documentation
|
| 116 |
+
files (the “Font Software”), to reproduce and distribute the Font Software,
|
| 117 |
+
including without limitation the rights to use, copy, merge, publish,
|
| 118 |
+
distribute,
|
| 119 |
+
and/or sell copies of the Font Software, and to permit persons to whom
|
| 120 |
+
the Font Software is furnished to do so, subject to the following
|
| 121 |
+
conditions:
|
| 122 |
+
|
| 123 |
+
The above copyright and trademark notices and this permission notice
|
| 124 |
+
shall be
|
| 125 |
+
included in all copies of one or more of the Font Software typefaces.
|
| 126 |
+
|
| 127 |
+
The Font Software may be modified, altered, or added to, and in particular
|
| 128 |
+
the designs of glyphs or characters in the Fonts may be modified and
|
| 129 |
+
additional
|
| 130 |
+
glyphs or characters may be added to the Fonts, only if the fonts are
|
| 131 |
+
renamed
|
| 132 |
+
to names not containing either the words “Bitstream” or the word “Vera”.
|
| 133 |
+
|
| 134 |
+
This License becomes null and void to the extent applicable to Fonts or
|
| 135 |
+
Font Software
|
| 136 |
+
that has been modified and is distributed under the “Bitstream Vera”
|
| 137 |
+
names.
|
| 138 |
+
|
| 139 |
+
The Font Software may be sold as part of a larger software package but
|
| 140 |
+
no copy
|
| 141 |
+
of one or more of the Font Software typefaces may be sold by itself.
|
| 142 |
+
|
| 143 |
+
THE FONT SOFTWARE IS PROVIDED “AS IS”, WITHOUT WARRANTY OF ANY KIND, EXPRESS
|
| 144 |
+
OR IMPLIED, INCLUDING BUT NOT LIMITED TO ANY WARRANTIES OF MERCHANTABILITY,
|
| 145 |
+
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT OF COPYRIGHT, PATENT,
|
| 146 |
+
TRADEMARK, OR OTHER RIGHT. IN NO EVENT SHALL BITSTREAM OR THE GNOME
|
| 147 |
+
FOUNDATION
|
| 148 |
+
BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER LIABILITY, INCLUDING ANY GENERAL,
|
| 149 |
+
SPECIAL, INDIRECT, INCIDENTAL, OR CONSEQUENTIAL DAMAGES, WHETHER IN AN
|
| 150 |
+
ACTION
|
| 151 |
+
OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF THE USE OR
|
| 152 |
+
INABILITY TO USE
|
| 153 |
+
THE FONT SOFTWARE OR FROM OTHER DEALINGS IN THE FONT SOFTWARE.
|
| 154 |
+
Except as contained in this notice, the names of GNOME, the GNOME
|
| 155 |
+
Foundation,
|
| 156 |
+
and Bitstream Inc., shall not be used in advertising or otherwise to promote
|
| 157 |
+
the sale, use or other dealings in this Font Software without prior written
|
| 158 |
+
authorization from the GNOME Foundation or Bitstream Inc., respectively.
|
| 159 |
+
For further information, contact: fonts at gnome dot org.
|
| 160 |
+
|
| 161 |
+
AMSFonts (v. 2.2) copyright
|
| 162 |
+
|
| 163 |
+
The PostScript Type 1 implementation of the AMSFonts produced by and
|
| 164 |
+
previously distributed by Blue Sky Research and Y&Y, Inc. are now freely
|
| 165 |
+
available for general use. This has been accomplished through the
|
| 166 |
+
cooperation
|
| 167 |
+
of a consortium of scientific publishers with Blue Sky Research and Y&Y.
|
| 168 |
+
Members of this consortium include:
|
| 169 |
+
|
| 170 |
+
Elsevier Science IBM Corporation Society for Industrial and Applied
|
| 171 |
+
Mathematics (SIAM) Springer-Verlag American Mathematical Society (AMS)
|
| 172 |
+
|
| 173 |
+
In order to assure the authenticity of these fonts, copyright will be
|
| 174 |
+
held by
|
| 175 |
+
the American Mathematical Society. This is not meant to restrict in any way
|
| 176 |
+
the legitimate use of the fonts, such as (but not limited to) electronic
|
| 177 |
+
distribution of documents containing these fonts, inclusion of these fonts
|
| 178 |
+
into other public domain or commercial font collections or computer
|
| 179 |
+
applications, use of the outline data to create derivative fonts and/or
|
| 180 |
+
faces, etc. However, the AMS does require that the AMS copyright notice be
|
| 181 |
+
removed from any derivative versions of the fonts which have been altered in
|
| 182 |
+
any way. In addition, to ensure the fidelity of TeX documents using Computer
|
| 183 |
+
Modern fonts, Professor Donald Knuth, creator of the Computer Modern faces,
|
| 184 |
+
has requested that any alterations which yield different font metrics be
|
| 185 |
+
given a different name.
|
| 186 |
+
|
| 187 |
+
$Id$
|
render_assets/abcjs-basic-min.js
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
render_assets/manifest.json
ADDED
|
@@ -0,0 +1,106 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
{
|
| 2 |
+
"abcjs_version": "6.6.3",
|
| 3 |
+
"abcjs_revision": "5dc285c5a4a918d33f2ab4b3c381499847e88426",
|
| 4 |
+
"piano_revision": "cbd6b6f6d1af89ebfb69402860741288f08ff8b7",
|
| 5 |
+
"piano_source": "https://github.com/paulrosen/midi-js-soundfonts/tree/cbd6b6f6d1af89ebfb69402860741288f08ff8b7/FluidR3_GM/acoustic_grand_piano-mp3",
|
| 6 |
+
"files": {
|
| 7 |
+
"DejaVuSans.ttf": "7da195a74c55bef988d0d48f9508bd5d849425c1770dba5d7bfc6ce9ed848954",
|
| 8 |
+
"LICENSE.abcjs": "ab1072f859916529e284a8f43643146de9935041bb347a35559a96c87ce5566f",
|
| 9 |
+
"LICENSE.font": "7a083b136e64d064794c3419751e5c7dd10d2f64c108fe5ba161eae5e5958a93",
|
| 10 |
+
"abcjs-basic-min.js": "abdab74cf95c39fb9ff4ae0c96735b9c35222851f0844ce471ddd4354739bc75",
|
| 11 |
+
"renderer.js": "0756398dfe69fe1c54c8779252885fbf573d02eecbe3f2389a0b81e0f43fcd09",
|
| 12 |
+
"soundfonts/ATTRIBUTION.md": "87237050093c1d4f47f5c459f1e18b57ea6710f92d1a57b4462f8f5aa31e11be",
|
| 13 |
+
"soundfonts/acoustic_grand_piano-mp3/A0.mp3": "5c47e8142b04505c888b8edd7246a8881b4b30c4340682f0264191599df232fc",
|
| 14 |
+
"soundfonts/acoustic_grand_piano-mp3/A1.mp3": "abab379d4d0557abb93856f816b7bb395872e656533c3fe7f02048fc993834b7",
|
| 15 |
+
"soundfonts/acoustic_grand_piano-mp3/A2.mp3": "e21d9ab9b30c73c0aea921c2cd0dd59fece4cc021619927c9ff3dee71f0a9679",
|
| 16 |
+
"soundfonts/acoustic_grand_piano-mp3/A3.mp3": "02dee2fa184b9b94a5f72886c2370bd076103ebb76fc179982fa503e549dcc10",
|
| 17 |
+
"soundfonts/acoustic_grand_piano-mp3/A4.mp3": "4f1671ac831bc41c845e67b18e6c9ea200fc9b753bdffe124703b7141aa928c3",
|
| 18 |
+
"soundfonts/acoustic_grand_piano-mp3/A5.mp3": "2cccfe0aa8077bfd20bcbba470012faa0328ae97d5bd2297147c01f2b255e574",
|
| 19 |
+
"soundfonts/acoustic_grand_piano-mp3/A6.mp3": "87dab2b5bfdd8bafa575fc24b99591e87e5a0c991a8bbdc8aeb8a1e517b576d2",
|
| 20 |
+
"soundfonts/acoustic_grand_piano-mp3/A7.mp3": "ace76663dac4610a93259655f069835e6eb397b76aedf85c1f48a5036cf25687",
|
| 21 |
+
"soundfonts/acoustic_grand_piano-mp3/Ab1.mp3": "a9889722caea82aa67a8b7f600f1def8e51e095f1f423949c3a8e9dbf5c14abd",
|
| 22 |
+
"soundfonts/acoustic_grand_piano-mp3/Ab2.mp3": "b45bb57e125cdbbc3bb521ae643e67a90c0ff646d37e5b5284e5c70bfeff6c75",
|
| 23 |
+
"soundfonts/acoustic_grand_piano-mp3/Ab3.mp3": "843d633061631861c943266b4836573fe8ba8b8a10df9bd94adc380376cdb8ce",
|
| 24 |
+
"soundfonts/acoustic_grand_piano-mp3/Ab4.mp3": "e8612f8602d1c146f9b7f69a441ee518876ccdfa469457366c8064c25d9d3116",
|
| 25 |
+
"soundfonts/acoustic_grand_piano-mp3/Ab5.mp3": "31c86a311eb68761b8008359b34addfeec249d1c7d6795a52f8f2e172e3e8a05",
|
| 26 |
+
"soundfonts/acoustic_grand_piano-mp3/Ab6.mp3": "7775c09d08d344af9528edbf7abb6f59ac600e4d15efd6f139845bb107a822ed",
|
| 27 |
+
"soundfonts/acoustic_grand_piano-mp3/Ab7.mp3": "66932305995ff04e9e0244c5a408ec87c9131ab3915eac15c72c7659132ab0f2",
|
| 28 |
+
"soundfonts/acoustic_grand_piano-mp3/B0.mp3": "4422d1d2f0498aec427165415f5d665d78699e8623475215ce9b27c18b3d331f",
|
| 29 |
+
"soundfonts/acoustic_grand_piano-mp3/B1.mp3": "048003df58298d2659bff084885ad69d9531fdc03837c7ee39c8bcf68334b5b9",
|
| 30 |
+
"soundfonts/acoustic_grand_piano-mp3/B2.mp3": "be2df872b00b68adead4ca24ef68bd4a93bf0c4864d455d742171e5df91e86b7",
|
| 31 |
+
"soundfonts/acoustic_grand_piano-mp3/B3.mp3": "82c0463af9c9fe51d44aa59b71d2dccefab11c1ea98e2748ee93244f8af9169e",
|
| 32 |
+
"soundfonts/acoustic_grand_piano-mp3/B4.mp3": "5f09187820410b0f78575f232c3509f367478b7755b6eb5aa46ce6cf4c16a0e8",
|
| 33 |
+
"soundfonts/acoustic_grand_piano-mp3/B5.mp3": "4e0932827d3e9390080c4a9acd420e9b83871a5f73723feaf61d2d7f6b628b01",
|
| 34 |
+
"soundfonts/acoustic_grand_piano-mp3/B6.mp3": "d31aa62152acd598d55fb02cc024e59e0176a7dc13d27813882f299cae7baf7a",
|
| 35 |
+
"soundfonts/acoustic_grand_piano-mp3/B7.mp3": "6e1aad1cc92f17ebf6fe6161766cce279250e67a30569649481c77ab80db1677",
|
| 36 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb0.mp3": "e4734f8d0285f0e5a39db7f61023510b46fb1570743b20d287479dcaa5536adb",
|
| 37 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb1.mp3": "6d4face1b6559a08af3f9a49fcbc472b43e974e072156b31850ab9a230713aa3",
|
| 38 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb2.mp3": "78cc11431a2a2d52294f0c714c59107a886f637806ae2cba63d50ff53585b93e",
|
| 39 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb3.mp3": "6ac0f67f18e15db35553d4fcd83ceec4473ebe639e31659ce93cf8466cb06644",
|
| 40 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb4.mp3": "d8f729899595116cb62ca9b7b801b4dfb9322bdf50d1bbfa847010e43a39f8f5",
|
| 41 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb5.mp3": "51224d4c26f1c89b08e43865ef55d3638367f7f40fbc92e4ab88610b54d4301d",
|
| 42 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb6.mp3": "394488d6cbd2c58b40f202e5acc238abdb73f261bcebd7105fb81d091113c4aa",
|
| 43 |
+
"soundfonts/acoustic_grand_piano-mp3/Bb7.mp3": "fe72a195f6036722e587e1597e54d66ec70e0ce11768e74e39bb2319bf256f80",
|
| 44 |
+
"soundfonts/acoustic_grand_piano-mp3/C1.mp3": "eb3b985e3b692cfe0c48550a26843fc3da44c7630c0b7bcfd820e87c5122cccc",
|
| 45 |
+
"soundfonts/acoustic_grand_piano-mp3/C2.mp3": "e02c5a80c03871e1889a3bb371d853f92294c1e8538ac06f70df0fd045540207",
|
| 46 |
+
"soundfonts/acoustic_grand_piano-mp3/C3.mp3": "7d0409ea3c40af25c6d3d606f12890f88f896d93c381b6eb8c58497f308b30d3",
|
| 47 |
+
"soundfonts/acoustic_grand_piano-mp3/C4.mp3": "b6d32e0d9d0c2e4b88165a2ec27b8787a970f131a03d96d065f28f158640af96",
|
| 48 |
+
"soundfonts/acoustic_grand_piano-mp3/C5.mp3": "7a4be6dae323ac5564ad1a617b554fd595e5e0a63f19c3581e3ee94c33ae2e4b",
|
| 49 |
+
"soundfonts/acoustic_grand_piano-mp3/C6.mp3": "6a0030c071ccd6c9a722455851d65b18cb86cbfe515772a6ded4c747eda339a0",
|
| 50 |
+
"soundfonts/acoustic_grand_piano-mp3/C7.mp3": "5e4cadc06eb25e2302dcc5bf1af76fe85ed258759d02c46e9169376ade83845e",
|
| 51 |
+
"soundfonts/acoustic_grand_piano-mp3/C8.mp3": "4ec066cab9362f6706b90fedc3249fa3c221769ee95ff52d7e46aaaf0a01fd16",
|
| 52 |
+
"soundfonts/acoustic_grand_piano-mp3/D1.mp3": "2389562397c7bdf2a856965b8b371eb6311be3cad2875c71d8ff4a607a242cea",
|
| 53 |
+
"soundfonts/acoustic_grand_piano-mp3/D2.mp3": "6233cc6f3670a287e33c7b8ece3b03a787ee1e169c8355432ce0b5e194bee612",
|
| 54 |
+
"soundfonts/acoustic_grand_piano-mp3/D3.mp3": "b8a5ed9cf147e10f484f3ad39930130d487b12045a35cef73f48dab401e6583b",
|
| 55 |
+
"soundfonts/acoustic_grand_piano-mp3/D4.mp3": "90539782bdbbbe714de579bfbbe6f478e07e112c6f78d03e8ee82462fbf2a569",
|
| 56 |
+
"soundfonts/acoustic_grand_piano-mp3/D5.mp3": "b25b5e541807af700f4e2272ac4c55a2f9b5cc8f176fcd9961d4456803af2b87",
|
| 57 |
+
"soundfonts/acoustic_grand_piano-mp3/D6.mp3": "46317b9dca22a75986ecfeba93ef34c8f7bc6d4e2e0223dc64b13c9a1984110b",
|
| 58 |
+
"soundfonts/acoustic_grand_piano-mp3/D7.mp3": "d7db5e63d7caa2fc62e00f55366f9a5236b3a13f7eb0d141a9823e10c7cc4209",
|
| 59 |
+
"soundfonts/acoustic_grand_piano-mp3/Db1.mp3": "84f2ba7505c515c39797eafe36295b15ee9da8adbafcfbfbfc7685532b4bdce4",
|
| 60 |
+
"soundfonts/acoustic_grand_piano-mp3/Db2.mp3": "c5d76b97326a189786aba25685c8789f59d4dc41d5fbd4d6a6aa97986b2b7d1d",
|
| 61 |
+
"soundfonts/acoustic_grand_piano-mp3/Db3.mp3": "c5bd77b92ad8411c47e509643a7d6fc9c74f3aaba2290e2f4aed6a2ba1446e88",
|
| 62 |
+
"soundfonts/acoustic_grand_piano-mp3/Db4.mp3": "f49892c4e935296db1cb785000818e6a04bf0756c4443bffa6245144fdaf3e47",
|
| 63 |
+
"soundfonts/acoustic_grand_piano-mp3/Db5.mp3": "924a70321fe728ff7a69c941029e060f529e6ada507f20357ee74039e270e588",
|
| 64 |
+
"soundfonts/acoustic_grand_piano-mp3/Db6.mp3": "5e4b55cc27db80e899ea817c546f4c0cf15b43267f426f8de53baac770f90be7",
|
| 65 |
+
"soundfonts/acoustic_grand_piano-mp3/Db7.mp3": "56c93ba30fa46db98590fd4c218fec88e87b02eaa0b177bcc0ad621fd9a506da",
|
| 66 |
+
"soundfonts/acoustic_grand_piano-mp3/Db8.mp3": "1be0e2193b729fc32a4fbb82f9f46bb4c0b2cbb67027c0f07c8eeff6460b64cf",
|
| 67 |
+
"soundfonts/acoustic_grand_piano-mp3/E1.mp3": "740bf03f7ac2a9bbd41ae10f7daf876bc8bde6fe6d355f2f2781ae70d95c762a",
|
| 68 |
+
"soundfonts/acoustic_grand_piano-mp3/E2.mp3": "316b392d51d9c27334976511643112b088a9812a8e37e598957e8acc6725aad0",
|
| 69 |
+
"soundfonts/acoustic_grand_piano-mp3/E3.mp3": "5e556a4850b93b1adba6655372c9e58913169864c52a8e6f232f4b1e901ca2dc",
|
| 70 |
+
"soundfonts/acoustic_grand_piano-mp3/E4.mp3": "48bcecf53832988af83f74aa65efd7daa34f1e3ebf10de4679363b89d54f3852",
|
| 71 |
+
"soundfonts/acoustic_grand_piano-mp3/E5.mp3": "c48cee1a11528da53e1f780a1ed36558530f053b5ed8c7033709b2002b452513",
|
| 72 |
+
"soundfonts/acoustic_grand_piano-mp3/E6.mp3": "e4de8378854d6b39b40d4cf7b022f60bda097993b492639c8a0a005f02a010e0",
|
| 73 |
+
"soundfonts/acoustic_grand_piano-mp3/E7.mp3": "e757f51d2fb54dca2c0adb4ee6a2b86459a52e2f7dfa227529558a6a0f28fd07",
|
| 74 |
+
"soundfonts/acoustic_grand_piano-mp3/Eb1.mp3": "694e89cd30008f47411336ea1b76675503bbc5383c3849df676b2cf6b51eee67",
|
| 75 |
+
"soundfonts/acoustic_grand_piano-mp3/Eb2.mp3": "708ff787e5ae2bfa45fe386e5fadc90a9339817dd245ed6e6f62532e42fad8d2",
|
| 76 |
+
"soundfonts/acoustic_grand_piano-mp3/Eb3.mp3": "bec9944789ef9e7ec3dd7f9fcdf63ed1e5e82e29447c81be3578df9b1269b4b6",
|
| 77 |
+
"soundfonts/acoustic_grand_piano-mp3/Eb4.mp3": "96586e2865651cfa0469e762f3bd40821f6e45eecf4761a86a37bc59cfebfb2c",
|
| 78 |
+
"soundfonts/acoustic_grand_piano-mp3/Eb5.mp3": "aad99b852c20c2aee66143dd103513523566716cb2f4c6ef6ac478ef84f282b0",
|
| 79 |
+
"soundfonts/acoustic_grand_piano-mp3/Eb6.mp3": "c0f4c45d7930ed90edebfd93321807cdb7346e364ba1933867c4b49ba18d5368",
|
| 80 |
+
"soundfonts/acoustic_grand_piano-mp3/Eb7.mp3": "c23b9b43ca8ece152e59ef04832b2290e96db54f9539af3162df4054c23271c7",
|
| 81 |
+
"soundfonts/acoustic_grand_piano-mp3/F1.mp3": "6b83be26ca990ee91e4d404afd5756e5d63b7a6f9ac46da068c80b2692945dba",
|
| 82 |
+
"soundfonts/acoustic_grand_piano-mp3/F2.mp3": "7bb9c26ec86e1d32194ab5ced3805e1026e638a4a5481d556a7e3f2860141c46",
|
| 83 |
+
"soundfonts/acoustic_grand_piano-mp3/F3.mp3": "ff912c3211f80dc0c4f541ebd39c1db11f21cfda684d54999f9bce87aaf0cf20",
|
| 84 |
+
"soundfonts/acoustic_grand_piano-mp3/F4.mp3": "a7426425ed95ea8652a888d850c34c4b766c7c7a11c168ed6899c38a6836be0c",
|
| 85 |
+
"soundfonts/acoustic_grand_piano-mp3/F5.mp3": "d7f046daec924ce8607ec92360e0d0e42e29e4b3a108839179e04aba4adadc51",
|
| 86 |
+
"soundfonts/acoustic_grand_piano-mp3/F6.mp3": "6f7794d298d1036ebbd514ad1b3605ffbd280556925fba05e29b83052cbf3dfd",
|
| 87 |
+
"soundfonts/acoustic_grand_piano-mp3/F7.mp3": "f49b2b54f0b913f636b5822eae373a85dbb30ec5e9a6ebc18e498fdfe0547011",
|
| 88 |
+
"soundfonts/acoustic_grand_piano-mp3/G1.mp3": "d187e74ca3a23d08248ea4073abc0893a758c8b15c4b6af16c3ba1d2394830db",
|
| 89 |
+
"soundfonts/acoustic_grand_piano-mp3/G2.mp3": "034f90b9f81e5aca0012f1ff73d731308653578a7eb4bee30efca37af5afce77",
|
| 90 |
+
"soundfonts/acoustic_grand_piano-mp3/G3.mp3": "3b75e4dfd657723b3c7a30bca0409dce8013b4083d54e703c8ec513b535cee99",
|
| 91 |
+
"soundfonts/acoustic_grand_piano-mp3/G4.mp3": "37d21e13d5cecc8c0f1645c83bfe82f6032407a3b7de3dde14021c196526664e",
|
| 92 |
+
"soundfonts/acoustic_grand_piano-mp3/G5.mp3": "2eb949282bc676f3dc96bf5008459b5fff02696e63b70ff799ff08c36bb97058",
|
| 93 |
+
"soundfonts/acoustic_grand_piano-mp3/G6.mp3": "c24ccf8d6f0c6c9818fe244d150a42c42585629db413e6b245088fdb589e9f3d",
|
| 94 |
+
"soundfonts/acoustic_grand_piano-mp3/G7.mp3": "7cc4630a06cbc11438fac85c00b640cc7677be8332105204632b687d7127ebc4",
|
| 95 |
+
"soundfonts/acoustic_grand_piano-mp3/Gb1.mp3": "6ff97bc16a3020037f46ebe4196ce2602eed8af243b567768d4cb5d47a8ba0bc",
|
| 96 |
+
"soundfonts/acoustic_grand_piano-mp3/Gb2.mp3": "95ab26a90d088fdc938673005ba2b2a4b166ee188ffbad0c7f75fb746334b043",
|
| 97 |
+
"soundfonts/acoustic_grand_piano-mp3/Gb3.mp3": "91605f55bf0265c1e526b966a3901f9c2813473ac890f4b7330ee08874c5d30f",
|
| 98 |
+
"soundfonts/acoustic_grand_piano-mp3/Gb4.mp3": "f0db7ee32b9175907ca60cf53e6c7a129ea8e7c8f6562dc7ebe6e5c81dcbdf99",
|
| 99 |
+
"soundfonts/acoustic_grand_piano-mp3/Gb5.mp3": "bd1d81203f5338ab62372f88c0b4eb5b4a3133ef58518a8770f3dfdf34f74124",
|
| 100 |
+
"soundfonts/acoustic_grand_piano-mp3/Gb6.mp3": "c5f87e3bf666919bd1e230f6945ee4af752761faa54dc1814f854e8e9d8a7920",
|
| 101 |
+
"soundfonts/acoustic_grand_piano-mp3/Gb7.mp3": "8615b92d93fa16b4051bdc8800a6430642bdcff17e26d020ec5fd83bf89db759"
|
| 102 |
+
},
|
| 103 |
+
"font_version": "DejaVu Sans 2.37",
|
| 104 |
+
"font_license_source": "https://raw.githubusercontent.com/dejavu-fonts/dejavu-fonts/version_2_37/LICENSE",
|
| 105 |
+
"abcjs_git_blob": "a56bbfc0d004c9b14ea17ad4fd6b52a61b74d255"
|
| 106 |
+
}
|
render_assets/renderer.js
ADDED
|
@@ -0,0 +1,165 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"use strict";
|
| 2 |
+
|
| 3 |
+
// These display adaptations leave the saved ABC unchanged.
|
| 4 |
+
function expandRests(abc) {
|
| 5 |
+
let header = true, voice = null, meter = [4, 4], unit = [1, 8];
|
| 6 |
+
const meters = new Map();
|
| 7 |
+
return abc.split(/\r?\n/).map(line => {
|
| 8 |
+
const field = line.match(/^([A-Za-z]):\s*(.*)$/);
|
| 9 |
+
if (field) {
|
| 10 |
+
const [, key, value] = field;
|
| 11 |
+
if (key === "V") {
|
| 12 |
+
voice = value.split(/\s/)[0];
|
| 13 |
+
if (!meters.has(voice)) meters.set(voice, meter);
|
| 14 |
+
}
|
| 15 |
+
if (key === "M") {
|
| 16 |
+
const match = value.match(/^(\d+)\/(\d+)$/);
|
| 17 |
+
if (match) {
|
| 18 |
+
const next = match.slice(1).map(Number);
|
| 19 |
+
if (header || !voice) meter = next;
|
| 20 |
+
else meters.set(voice, next);
|
| 21 |
+
}
|
| 22 |
+
}
|
| 23 |
+
if (key === "L") unit = value.split("/").map(Number);
|
| 24 |
+
if (key === "K") header = false;
|
| 25 |
+
return line;
|
| 26 |
+
}
|
| 27 |
+
if (header || line.startsWith("%")) return line;
|
| 28 |
+
return line.replace(/"[^"\n]*"|\[[^\]\n]*\]|Z(\d*)\|/g, (token, count) => {
|
| 29 |
+
if (count === undefined) return token;
|
| 30 |
+
const [n, d] = meters.get(voice) || meter;
|
| 31 |
+
return (`z${n * unit[1] / (d * unit[0])}|`).repeat(Number(count || 1));
|
| 32 |
+
});
|
| 33 |
+
}).join("\n");
|
| 34 |
+
}
|
| 35 |
+
|
| 36 |
+
function fixSystemKeys(tune) {
|
| 37 |
+
for (const line of tune.lines) for (const staff of line.staff || []) {
|
| 38 |
+
for (const voice of staff.voices) {
|
| 39 |
+
for (let i = 0; i < voice.length && voice[i].el_type !== "note";) {
|
| 40 |
+
const event = voice[i];
|
| 41 |
+
if (event.el_type === "key" || event.el_type === "keySignature") {
|
| 42 |
+
staff.key = {...event, accidentals: event.accidentals.map(a => ({...a}))};
|
| 43 |
+
delete staff.key.impliedNaturals;
|
| 44 |
+
voice.splice(i, 1);
|
| 45 |
+
} else i++;
|
| 46 |
+
}
|
| 47 |
+
}
|
| 48 |
+
}
|
| 49 |
+
}
|
| 50 |
+
|
| 51 |
+
window.renderScore = async abc => {
|
| 52 |
+
document.body.innerHTML = '<div id="source"></div><main id="pages"></main>';
|
| 53 |
+
const fontStyle = document.createElement("style");
|
| 54 |
+
fontStyle.textContent = '#source text, #source span, #pages text {font-family:SheetSageSans !important}';
|
| 55 |
+
document.head.append(fontStyle);
|
| 56 |
+
const source = document.getElementById("source");
|
| 57 |
+
source.style.cssText = "position:absolute;left:-2000px;width:714px";
|
| 58 |
+
const tunes = ABCJS.renderAbc(source, expandRests(abc), {
|
| 59 |
+
staffwidth: 674, paddingtop: 12, paddingbottom: 12, paddingleft: 20, paddingright: 20,
|
| 60 |
+
oneSvgPerLine: true, print: true, add_classes: true,
|
| 61 |
+
wrap: {minSpacing: 1.8, maxSpacing: 2.7, preferredMeasuresPerLine: 4},
|
| 62 |
+
afterParsing: fixSystemKeys,
|
| 63 |
+
});
|
| 64 |
+
if (!tunes.length || !tunes.some(t => t.lines.some(l => l.staff?.length)))
|
| 65 |
+
throw new Error("ABC contains no music to render.");
|
| 66 |
+
const warnings = tunes.flatMap(t => t.warnings || []);
|
| 67 |
+
if (warnings.length) throw new Error("Invalid ABC: " + warnings.join("; "));
|
| 68 |
+
await document.fonts.ready;
|
| 69 |
+
const ns = "http://www.w3.org/2000/svg";
|
| 70 |
+
const pages = [];
|
| 71 |
+
let page, content, used = 0, systems = 0;
|
| 72 |
+
function addPage() {
|
| 73 |
+
page = document.createElementNS(ns, "svg");
|
| 74 |
+
for (const [k, v] of Object.entries({xmlns:ns, width:794, height:1123, viewBox:"0 0 794 1123"}))
|
| 75 |
+
page.setAttribute(k,v);
|
| 76 |
+
page.setAttribute("class", "score-page");
|
| 77 |
+
const fontStyle = document.createElementNS(ns, "style");
|
| 78 |
+
fontStyle.textContent = '@font-face{font-family:SheetSageSans;src:url(data:font/ttf;base64,' + window.renderFontData + ')}text{font-family:SheetSageSans !important}';
|
| 79 |
+
page.append(fontStyle);
|
| 80 |
+
const license = document.createElementNS(ns, "metadata");
|
| 81 |
+
license.textContent = window.renderFontLicense;
|
| 82 |
+
page.append(license);
|
| 83 |
+
const background = document.createElementNS(ns, "rect");
|
| 84 |
+
for (const [k,v] of Object.entries({width:794,height:1123,fill:"white"})) background.setAttribute(k,v);
|
| 85 |
+
page.append(background);
|
| 86 |
+
content = document.createElementNS(ns, "g");
|
| 87 |
+
content.setAttribute("transform", "translate(40 40)");
|
| 88 |
+
page.append(content);
|
| 89 |
+
document.getElementById("pages").append(page);
|
| 90 |
+
pages.push(page);
|
| 91 |
+
used = 0;
|
| 92 |
+
}
|
| 93 |
+
for (const svg of source.querySelectorAll("svg")) {
|
| 94 |
+
const box = svg.getBoundingClientRect();
|
| 95 |
+
if (!(box.width > 0 && box.height > 0)) continue;
|
| 96 |
+
const scale = Math.min(1, 714 / box.width);
|
| 97 |
+
const height = box.height * scale;
|
| 98 |
+
if (height > 1043) throw new Error("A staff system is taller than one page.");
|
| 99 |
+
if (!page || used + height > 1043) addPage();
|
| 100 |
+
const group = document.createElementNS(ns, "g");
|
| 101 |
+
group.setAttribute("transform", `translate(0 ${used}) scale(${scale})`);
|
| 102 |
+
const clone = svg.cloneNode(true);
|
| 103 |
+
clone.setAttribute("width",box.width);
|
| 104 |
+
clone.setAttribute("height",box.height);
|
| 105 |
+
clone.style.overflow = "visible";
|
| 106 |
+
group.append(clone);
|
| 107 |
+
content.append(group);
|
| 108 |
+
used += height;
|
| 109 |
+
systems++;
|
| 110 |
+
}
|
| 111 |
+
source.remove();
|
| 112 |
+
if (!pages.length) throw new Error("No score pages were produced.");
|
| 113 |
+
return {pages:pages.map(p => new XMLSerializer().serializeToString(p)), systems, warnings};
|
| 114 |
+
};
|
| 115 |
+
|
| 116 |
+
window.renderAudio = async ({tracks, duration}) => {
|
| 117 |
+
const context = new AudioContext({sampleRate:44100});
|
| 118 |
+
await context.resume();
|
| 119 |
+
const sequence = new ABCJS.synth.SynthSequence();
|
| 120 |
+
for (const track of tracks) {
|
| 121 |
+
const id = sequence.addTrack();
|
| 122 |
+
sequence.setInstrument(id, 0);
|
| 123 |
+
for (const note of track.notes) sequence.tracks[id].push({
|
| 124 |
+
cmd:"note", instrument:0, pitch:note.pitch, volume:note.velocity,
|
| 125 |
+
start:note.start, duration:note.end-note.start, gap:0,
|
| 126 |
+
});
|
| 127 |
+
}
|
| 128 |
+
sequence.totalDuration = duration;
|
| 129 |
+
const synth = new ABCJS.synth.CreateSynth();
|
| 130 |
+
const loaded = await synth.init({audioContext:context, sequence, millisecondsPerMeasure:1000,
|
| 131 |
+
options:{soundFontUrl:"https://render.invalid/soundfonts/", fadeLength:35,
|
| 132 |
+
programOffsets:{acoustic_grand_piano:0}, soundFontVolumeMultiplier:1}});
|
| 133 |
+
if (loaded.error?.length) throw new Error("Could not load piano samples: " + loaded.error.join(", "));
|
| 134 |
+
await synth.prime();
|
| 135 |
+
const buffer = synth.getAudioBuffer();
|
| 136 |
+
if (!buffer) throw new Error("The renderer returned no audio.");
|
| 137 |
+
// Simultaneous voices retain their balance; a peak guard avoids PCM clipping.
|
| 138 |
+
let peak = 0;
|
| 139 |
+
for (let c=0; c<buffer.numberOfChannels; c++) {
|
| 140 |
+
const channel = buffer.getChannelData(c);
|
| 141 |
+
for (let i=0; i<channel.length; i++) peak = Math.max(peak, Math.abs(channel[i]));
|
| 142 |
+
}
|
| 143 |
+
const gain = Math.min(0.7, peak > 0 ? 0.98 / peak : 0.7);
|
| 144 |
+
window.renderedAudio = {buffer, gain};
|
| 145 |
+
await context.close();
|
| 146 |
+
return {frames:buffer.length, sample_rate:buffer.sampleRate, channels:buffer.numberOfChannels, gain};
|
| 147 |
+
};
|
| 148 |
+
|
| 149 |
+
window.audioBlock = ({start, count}) => {
|
| 150 |
+
const {buffer, gain} = window.renderedAudio;
|
| 151 |
+
const frames = Math.min(count, buffer.length-start);
|
| 152 |
+
const bytes = new Uint8Array(frames*buffer.numberOfChannels*2);
|
| 153 |
+
const view = new DataView(bytes.buffer);
|
| 154 |
+
for (let c=0; c<buffer.numberOfChannels; c++) {
|
| 155 |
+
const channel = buffer.getChannelData(c);
|
| 156 |
+
for (let i=0; i<frames; i++) {
|
| 157 |
+
const sample = Math.max(-1,Math.min(1,channel[start+i]*gain));
|
| 158 |
+
view.setInt16((i*buffer.numberOfChannels+c)*2, Math.round(sample*32767), true);
|
| 159 |
+
}
|
| 160 |
+
}
|
| 161 |
+
let binary="";
|
| 162 |
+
for (let offset=0; offset<bytes.length; offset+=8192)
|
| 163 |
+
binary += String.fromCharCode(...bytes.subarray(offset,offset+8192));
|
| 164 |
+
return btoa(binary);
|
| 165 |
+
};
|
render_assets/soundfonts/ATTRIBUTION.md
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
# Piano samples
|
| 2 |
+
|
| 3 |
+
FluidR3 GM by Frank Wen. License: Creative Commons Attribution 3.0 US.
|
| 4 |
+
https://creativecommons.org/licenses/by/3.0/us/
|
| 5 |
+
|
| 6 |
+
Pre-rendered MP3 samples from Paul Rosen's fork of gleitz/midi-js-soundfonts:
|
| 7 |
+
https://github.com/paulrosen/midi-js-soundfonts
|
| 8 |
+
https://raw.githubusercontent.com/paulrosen/midi-js-soundfonts/gh-pages/FluidR3_GM/acoustic_grand_piano-mp3.js
|
| 9 |
+
|
| 10 |
+
The base64 samples were unpacked into individual MP3 files without changing the
|
| 11 |
+
audio, for same-origin/offline ABCJS score playback. Only piano is bundled.
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A0.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A1.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A2.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A3.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A4.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A5.mp3
ADDED
|
Binary file (18.9 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A6.mp3
ADDED
|
Binary file (15.3 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/A7.mp3
ADDED
|
Binary file (14.5 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/Ab1.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/Ab2.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/Ab3.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/Ab4.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/Ab5.mp3
ADDED
|
Binary file (20.4 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/Ab6.mp3
ADDED
|
Binary file (15.8 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/Ab7.mp3
ADDED
|
Binary file (14.7 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/B0.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|
render_assets/soundfonts/acoustic_grand_piano-mp3/B1.mp3
ADDED
|
Binary file (25.6 kB). View file
|
|
|