rAVEUK
/

rAVEUK a43992899 commited on
Commit
cecd352
·
0 Parent(s):

Duplicate from m-a-p/SheetSage2

Browse files

Co-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
Files changed (50) hide show
  1. .gitattributes +3 -0
  2. LICENSE +424 -0
  3. README.md +201 -0
  4. THIRD_PARTY_NOTICES.md +8 -0
  5. __init__.py +1 -0
  6. assets/architecture.png +3 -0
  7. audio_sheetsage2.py +92 -0
  8. benchmark_results.json +379 -0
  9. config.json +75 -0
  10. configuration_mert2.py +76 -0
  11. configuration_sheetsage2.py +88 -0
  12. durations_sheetsage2.py +14 -0
  13. exports_sheetsage2.py +200 -0
  14. generation_sheetsage2.py +634 -0
  15. infer.py +78 -0
  16. io_sheetsage2.py +67 -0
  17. labels_sheetsage2.py +29 -0
  18. midi_sheetsage2.py +92 -0
  19. model.safetensors +3 -0
  20. modeling_mert2.py +361 -0
  21. modeling_sheetsage2.py +448 -0
  22. notation_sheetsage2.py +1570 -0
  23. pipeline_sheetsage2.py +259 -0
  24. processing_sheetsage2.py +109 -0
  25. processor_config.json +11 -0
  26. render.py +45 -0
  27. render_assets/DejaVuSans.ttf +3 -0
  28. render_assets/LICENSE.abcjs +21 -0
  29. render_assets/LICENSE.font +187 -0
  30. render_assets/abcjs-basic-min.js +0 -0
  31. render_assets/manifest.json +106 -0
  32. render_assets/renderer.js +165 -0
  33. render_assets/soundfonts/ATTRIBUTION.md +11 -0
  34. render_assets/soundfonts/acoustic_grand_piano-mp3/A0.mp3 +0 -0
  35. render_assets/soundfonts/acoustic_grand_piano-mp3/A1.mp3 +0 -0
  36. render_assets/soundfonts/acoustic_grand_piano-mp3/A2.mp3 +0 -0
  37. render_assets/soundfonts/acoustic_grand_piano-mp3/A3.mp3 +0 -0
  38. render_assets/soundfonts/acoustic_grand_piano-mp3/A4.mp3 +0 -0
  39. render_assets/soundfonts/acoustic_grand_piano-mp3/A5.mp3 +0 -0
  40. render_assets/soundfonts/acoustic_grand_piano-mp3/A6.mp3 +0 -0
  41. render_assets/soundfonts/acoustic_grand_piano-mp3/A7.mp3 +0 -0
  42. render_assets/soundfonts/acoustic_grand_piano-mp3/Ab1.mp3 +0 -0
  43. render_assets/soundfonts/acoustic_grand_piano-mp3/Ab2.mp3 +0 -0
  44. render_assets/soundfonts/acoustic_grand_piano-mp3/Ab3.mp3 +0 -0
  45. render_assets/soundfonts/acoustic_grand_piano-mp3/Ab4.mp3 +0 -0
  46. render_assets/soundfonts/acoustic_grand_piano-mp3/Ab5.mp3 +0 -0
  47. render_assets/soundfonts/acoustic_grand_piano-mp3/Ab6.mp3 +0 -0
  48. render_assets/soundfonts/acoustic_grand_piano-mp3/Ab7.mp3 +0 -0
  49. render_assets/soundfonts/acoustic_grand_piano-mp3/B0.mp3 +0 -0
  50. 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/">🎵&nbsp;YuE2&nbsp;project</a>
20
+ ·
21
+ <a href="#quick-start">🚀&nbsp;Quick&nbsp;start</a>
22
+ ·
23
+ <a href="#benchmarks">📊&nbsp;Benchmarks</a>
24
+ ·
25
+ <a href="#citation">📚&nbsp;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&amp;logoColor=FFD21E" height="20" /></a>
29
+ &nbsp;
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&amp;logoColor=FFD21E" height="20" /></a>
31
+ &nbsp;
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&amp;logoColor=FFD21E" height="20" /></a>
33
+ &nbsp;
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&amp;logoColor=FFD21E" height="20" /></a>
35
+ &nbsp;
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&amp;logoColor=FFD21E" height="20" /></a>
37
+ &nbsp;
38
+ <a href="https://huggingface.co/datasets/m-a-p/WildSongBench"><img alt="🤗 WildSongBench" src="https://img.shields.io/badge/WildSongBench-374151?logo=huggingface&amp;logoColor=FFD21E" height="20" /></a>
39
+ &nbsp;
40
+ <a href="https://huggingface.co/m-a-p/SheetSage2"><img alt="SheetSage2" src="https://img.shields.io/badge/SheetSage2-374151?logo=huggingface&amp;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
+ ![SheetSage2 architecture](assets/architecture.png)
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

  • SHA256: 04e62e9aae6606edab1c81e5bbdf5cdba29786841407e6b674f4e06f132e364b
  • Pointer size: 131 Bytes
  • Size of remote file: 339 kB
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