From 62a1cd5cec2fbb1b6be30595641e60066730c7e4 Mon Sep 17 00:00:00 2001 From: BobRandomNumber Date: Tue, 8 Jul 2025 22:18:02 -0400 Subject: [PATCH] Initial --- CITATION.cff | 9 + LICENSE | 23 + NOTICE | 30 + README.md | 81 ++ __init__.py | 3 + example_workflows/KyutaiTTS.json | 101 +++ example_workflows/KyutaiTTS.png | Bin 0 -> 90643 bytes moshi_src/LICENSE | 23 + moshi_src/LICENSE.audiocraft | 21 + moshi_src/moshi/__init__.py | 19 + moshi_src/moshi/conditioners/__init__.py | 10 + moshi_src/moshi/conditioners/base.py | 432 ++++++++++ moshi_src/moshi/conditioners/tensors.py | 16 + moshi_src/moshi/conditioners/text.py | 134 ++++ moshi_src/moshi/models/__init__.py | 14 + moshi_src/moshi/models/compression.py | 488 ++++++++++++ moshi_src/moshi/models/lm.py | 837 ++++++++++++++++++++ moshi_src/moshi/models/lm_utils.py | 124 +++ moshi_src/moshi/models/loaders.py | 481 ++++++++++++ moshi_src/moshi/models/tts.py | 621 +++++++++++++++ moshi_src/moshi/modules/__init__.py | 23 + moshi_src/moshi/modules/conv.py | 423 ++++++++++ moshi_src/moshi/modules/conv_test.py | 157 ++++ moshi_src/moshi/modules/gating.py | 115 +++ moshi_src/moshi/modules/lora.py | 122 +++ moshi_src/moshi/modules/resample.py | 119 +++ moshi_src/moshi/modules/rope.py | 90 +++ moshi_src/moshi/modules/seanet.py | 392 ++++++++++ moshi_src/moshi/modules/seanet_test.py | 187 +++++ moshi_src/moshi/modules/streaming.py | 217 +++++ moshi_src/moshi/modules/transformer.py | 956 +++++++++++++++++++++++ moshi_src/moshi/quantization/__init__.py | 13 + moshi_src/moshi/quantization/base.py | 170 ++++ moshi_src/moshi/quantization/core_vq.py | 528 +++++++++++++ moshi_src/moshi/quantization/vq.py | 318 ++++++++ moshi_src/moshi/utils/__init__.py | 10 + moshi_src/moshi/utils/autocast.py | 45 ++ moshi_src/moshi/utils/compile.py | 287 +++++++ moshi_src/moshi/utils/quantize.py | 57 ++ moshi_src/moshi/utils/sampling.py | 127 +++ moshi_src/moshi/utils/utils.py | 52 ++ nodes.py | 184 +++++ requirements.txt | 1 + 43 files changed, 8060 insertions(+) create mode 100644 CITATION.cff create mode 100644 LICENSE create mode 100644 NOTICE create mode 100644 README.md create mode 100644 __init__.py create mode 100644 example_workflows/KyutaiTTS.json create mode 100644 example_workflows/KyutaiTTS.png create mode 100644 moshi_src/LICENSE create mode 100644 moshi_src/LICENSE.audiocraft create mode 100644 moshi_src/moshi/__init__.py create mode 100644 moshi_src/moshi/conditioners/__init__.py create mode 100644 moshi_src/moshi/conditioners/base.py create mode 100644 moshi_src/moshi/conditioners/tensors.py create mode 100644 moshi_src/moshi/conditioners/text.py create mode 100644 moshi_src/moshi/models/__init__.py create mode 100644 moshi_src/moshi/models/compression.py create mode 100644 moshi_src/moshi/models/lm.py create mode 100644 moshi_src/moshi/models/lm_utils.py create mode 100644 moshi_src/moshi/models/loaders.py create mode 100644 moshi_src/moshi/models/tts.py create mode 100644 moshi_src/moshi/modules/__init__.py create mode 100644 moshi_src/moshi/modules/conv.py create mode 100644 moshi_src/moshi/modules/conv_test.py create mode 100644 moshi_src/moshi/modules/gating.py create mode 100644 moshi_src/moshi/modules/lora.py create mode 100644 moshi_src/moshi/modules/resample.py create mode 100644 moshi_src/moshi/modules/rope.py create mode 100644 moshi_src/moshi/modules/seanet.py create mode 100644 moshi_src/moshi/modules/seanet_test.py create mode 100644 moshi_src/moshi/modules/streaming.py create mode 100644 moshi_src/moshi/modules/transformer.py create mode 100644 moshi_src/moshi/quantization/__init__.py create mode 100644 moshi_src/moshi/quantization/base.py create mode 100644 moshi_src/moshi/quantization/core_vq.py create mode 100644 moshi_src/moshi/quantization/vq.py create mode 100644 moshi_src/moshi/utils/__init__.py create mode 100644 moshi_src/moshi/utils/autocast.py create mode 100644 moshi_src/moshi/utils/compile.py create mode 100644 moshi_src/moshi/utils/quantize.py create mode 100644 moshi_src/moshi/utils/sampling.py create mode 100644 moshi_src/moshi/utils/utils.py create mode 100644 nodes.py create mode 100644 requirements.txt diff --git a/CITATION.cff b/CITATION.cff new file mode 100644 index 0000000..fedba81 --- /dev/null +++ b/CITATION.cff @@ -0,0 +1,9 @@ +@techreport{kyutai2024moshi, + author = {Alexandre D\'efossez and Laurent Mazar\'e and Manu Orsini and Am\'elie Royer and + Patrick P\'erez and Herv\'e J\'egou and Edouard Grave and Neil Zeghidour}, + title = {Moshi: a speech-text foundation model for real-time dialogue}, + institution = {Kyutai}, + year={2024}, + month={September}, + url={http://kyutai.org/Moshi.pdf}, +} \ No newline at end of file diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..31aa793 --- /dev/null +++ b/LICENSE @@ -0,0 +1,23 @@ +Permission is hereby granted, free of charge, to any +person obtaining a copy of this software and associated +documentation files (the "Software"), to deal in the +Software without restriction, including without +limitation the rights to use, copy, modify, merge, +publish, distribute, sublicense, and/or sell copies of +the Software, and to permit persons to whom the Software +is furnished to do so, subject to the following +conditions: + +The above copyright notice and this permission notice +shall be included in all copies or substantial portions +of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF +ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED +TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A +PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT +SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY +CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION +OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR +IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER +DEALINGS IN THE SOFTWARE. diff --git a/NOTICE b/NOTICE new file mode 100644 index 0000000..d186c16 --- /dev/null +++ b/NOTICE @@ -0,0 +1,30 @@ +This software contains code from the 'moshi' project, developed by Kyutai. + +The original source code can be found at: +https://github.com/kyutai-labs/moshi/tree/main/moshi + +The 'moshi' project is licensed under the MIT License, the text of which is included below. + +-------------------------------------------------------------------------------- + +MIT License + +Copyright (c) 2024 Kyutai + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md new file mode 100644 index 0000000..708f4a3 --- /dev/null +++ b/README.md @@ -0,0 +1,81 @@ +# ComfyUI-KyutaiTTS + + +A custom node for ComfyUI that allows TTS generation with the [Kyutai TTS 1.6b en_fr model](https://huggingface.co/kyutai/tts-1.6b-en_fr) using [Kyutai offered voice models.](https://huggingface.co/kyutai/tts-voices) +The model's intended use is [https://github.com/kyutai-labs/delayed-streams-modeling](https://github.com/kyutai-labs/delayed-streams-modeling)which is not implemented here. +I made this version as it can generate large ammounts quickly and at acceptable quality for my use cases. +The model outputs at 24000Hz, some post processing can improve it if needed. + +## Features + +* **Text-to-Speech Generation:** Convert input text into spoken audio. +* **Model Selection:** Enter path to the Kyutai TTS model directory. +* **Voice Model Support:** Utilize various voice models (checkpoints) to customize the generated voice. +* **Device Selection:** Choose between CPU and CUDA (GPU) for processing, leveraging GPU acceleration for faster generation. +* **Adjustable Parameters:** Fine-tune speech generation with parameters such as: + * `n_q`: Number of quantization levels. + * `temp`: Sampling temperature for speech variability. + * `cfg_coef`: Classifier-free guidance coefficient for adherence to input. + * `padding_between`: Control silence between speech segments. + * `seed`: For reproducible audio generation. + +## Installation + +1. **Clone this Repository and install requirements:** + ```bash + cd ComfyUI/custom_nodes + git clone https://github.com/BobRandomNumber/ComfyUI-KyutaiTTS.git + pip install -r requirements.txt + ``` + +2. **Download Model Files:** + + * **Kyutai TTS Model (1.6B en_fr):** + Download all files from the `main` branch of the Hugging Face repository: + [https://huggingface.co/kyutai/tts-1.6b-en_fr/tree/main](https://huggingface.co/kyutai/tts-1.6b-en_fr/tree/main) + Create a dedicated folder for these files within your ComfyUI setup (e.g., `ComfyUI/models/checkpoints/KyutaiTTS`). The node expects the following files within this folder: + * `dsm_tts_1e68beda@240.safetensors` (Moshi weights) + * `tokenizer-e351c8d8-checkpoint125.safetensors` (Mimi weights) + * `tokenizer_spm_8k_en_fr_audio.model` (Tokenizer model) + * `config.json` (Model configuration) + + * **Kyutai TTS Voice Models:** + Download your desired voice models from the Hugging Face repository: + [https://huggingface.co/kyutai/tts-voices/tree/main](https://huggingface.co/kyutai/tts-voices/tree/main) + Place these voice models into your `ComfyUI/models/loras` directory or a subdirectory in loras. + +## Usage + +1. **Start ComfyUI:** +2. **Add the Node:** In your ComfyUI workflow, right-click and navigate to `Add Node` -> `Kyutai` -> `KyutaiTTS`. +3. **Set parameters:** + * **`text`**: Input the text you want to convert to speech. + * **`model_path`**: Input the path to the directory where you placed your Kyutai TTS model files (e.g., `C:\ComfyUI\models\checkpoints\KyutaiTTS`). + * **`voice_model`**: Select your desired voice model from the dropdown list. + * **`device`**: Choose `cuda` for GPU acceleration (recommended if available) or `cpu`. + * Adjust other parameters (`n_q`, `temp`, `cfg_coef`, `padding_between`, `seed`) as needed. +4. **Connect Output:** Connect the `AUDIO` output of the `KyutaiTTS` node to an audio playback or save node. +5. **Queue Prompt:** Queue your prompt to generate the audio. + +## Example Workflow + +An example workflow is provided in the `example_workflows` directory. + +![ComfyUI-KyutaiTTS Workflow Example](example_workflows/KyutaiTTS.png) + +## Troubleshooting + +* **`FileNotFoundError`**: Ensure that the `model_path` selected in the node points to the correct directory containing all required model files. Also, verify that your voice models are correctly placed and accessible by ComfyUI. +* **General Issues**: Always restart ComfyUI after making changes to custom nodes or model paths. Check the ComfyUI console for any error messages or warnings. + +## License + +This node pack integrates the Kyutai TTS model. Please refer to the original Kyutai TTS project's licensing information for details regarding the model and its components. + +## Attribution + +This custom node utilizes the Kyutai TTS model. + +* **Original Moshi Source:** [https://github.com/kyutai-labs/moshi/tree/main/moshi](https://github.com/kyutai-labs/moshi/tree/main/moshi) +* **Kyutai TTS Model (1.6B en_fr):** [https://huggingface.co/kyutai/tts-1.6b-en_fr](https://huggingface.co/kyutai/tts-1.6b-en_fr) +* **Kyutai TTS Voice Models:** [https://huggingface.co/kyutai/tts-voices](https://huggingface.co/kyutai/tts-voices) \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..d721463 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS'] \ No newline at end of file diff --git a/example_workflows/KyutaiTTS.json b/example_workflows/KyutaiTTS.json new file mode 100644 index 0000000..afd39d4 --- /dev/null +++ b/example_workflows/KyutaiTTS.json @@ -0,0 +1,101 @@ +{ + "id": "00000000-0000-0000-0000-000000000000", + "revision": 0, + "last_node_id": 3, + "last_link_id": 1, + "nodes": [ + { + "id": 3, + "type": "SaveAudio", + "pos": [ + 1370, + 370 + ], + "size": [ + 320, + 112 + ], + "flags": {}, + "order": 1, + "mode": 0, + "inputs": [ + { + "name": "audio", + "type": "AUDIO", + "link": 1 + } + ], + "outputs": [], + "properties": { + "cnr_id": "comfy-core", + "ver": "0.3.43" + }, + "widgets_values": [ + "audio/KyutaiTTS" + ] + }, + { + "id": 2, + "type": "KyutaiTTS", + "pos": [ + 670, + 370 + ], + "size": [ + 670, + 304 + ], + "flags": {}, + "order": 0, + "mode": 0, + "inputs": [], + "outputs": [ + { + "name": "AUDIO", + "type": "AUDIO", + "links": [ + 1 + ] + } + ], + "properties": { + "Node name for S&R": "KyutaiTTS" + }, + "widgets_values": [ + "Hey there! How are you?", + "C:\\AI\\ComfyUI\\models\\checkpoints\\KyutaiTTS", + "tts-voices\\expresso\\ex03-ex01_laughing_002_channel2_232s.wav.1e68beda@240.safetensors", + "cuda", + 32, + 0.6, + 2, + 1, + 1650465257, + "randomize" + ] + } + ], + "links": [ + [ + 1, + 2, + 0, + 3, + 0, + "AUDIO" + ] + ], + "groups": [], + "config": {}, + "extra": { + "ds": { + "scale": 1.2100000000000002, + "offset": [ + -525.0992340844525, + -107.91739572572313 + ] + }, + "frontendVersion": "1.23.4" + }, + "version": 0.4 +} \ No newline at end of file diff --git a/example_workflows/KyutaiTTS.png b/example_workflows/KyutaiTTS.png new file mode 100644 index 0000000000000000000000000000000000000000..03aa4c6786b48a537e594ba0a7502b2443ba6fb8 GIT binary patch literal 90643 zcmeEu2{@GN8@J>{sZgO3gHBWkDT)|{RQ8a{GTF;E7{)TjB$Z02qmV*rv4!mW5+T`A zWH%w(Fb0zujG6hKcj%mS&i{YD<+{G>`mXCc*LBYEzVp7%dO!F5`~B{F=C+}}_NEOx zH?Xm>Z908Q^DGLt{_2^?nvkfc?BgoY4ES5v^eUDKX}Gm;h>zfy!=6M z!@>uZAteXiI! z+IxX<|EE{Le3rX*^LBCh$KF}vH$DAVuYcY6c@7BHBY*4f$HS}p0Sk)s_F7ppaN+KW zbhr2Px&mz!M%>oTleKx`wn$e8KN(x3r#YFPaY8dF9 zfNc41M)32~|K6b;IN!F=j;~nxpV$A^&B}xS-h&&$VUZfAHUz?APc2y=#=cz1?5^tz7;0PFN4EZrSo@*YXd}^oPS;F|Hq%EB*Yq zgpMF9Bv@iWgR9+>J~)4((CeZONKa8?iL<~h{pDA{X#Vb*PTBj3dR?~nv=Pv(V^VI3%N z-`*F51QZGh{UxU$1O6kAaIyAwybNujoZLZ#?PY5>H+z?Zh=U3TQ3rgi(Ff%16_421 z+gYn0JS2AjW$j?^W$%VUdZNHsY`yKQ#ibPvO3NKkls+geFDc6e@_gYS))4%g=B%HP}d>;Rfn_#P- zDAuo6TBA)@+hPPClBr7;TW_Oz?V7ltR?#Log+}@(g+6@7;d3p7!x+*WBaNC-M!8rK zS0|}V|B{-DrrQM8i z?TP=Bt5+o>w*=?v>grzZl$4Z=J9{#Z8rD^Y@%i%V_U-k?JvIcxhGa_lfYXlla((Zw zDC1ZZ14uF*~j`#IauDU`<1c5_6Enh9~{sG6)$hiz>^%t86A zv?wFuL>tc+_o!2N#hZq1n;&rAF%|Yriw$-vF~JDjk=`cr+(OqP#oh!{yF~{5 zs72;}b3k)|TWDlzxY@U&CY)*@-6&LL)A3a!_a<h1K@IVD)B^{cAIQNm!0^*t+(p=>zAHwmPl2#K`028emm-CbyCzp z+6Y<^%Xg-Euy1}B9>C^H2pZtk$w5?r)yb!+S!ww?c?J(z%kqWzx~IYOn8X*`5AWV( zTs)bVX$=1`LcV{_kqjGlIVhD!0wXWnPxh0irgif`^Vp9|W{servE?-Srb0}NMe8o3 ziYBRA0-ITTP(eLiw51-*TCbUh^{)K=W3r*1n)IK>f0oGS%jf%JGib`gPrH5E=d)0i-J$wlEDL$)DWOvGW$h?i`f){!f^^;~pc1otI zQk|;dKRWYvHze-Ry@`{1JBo(YzJbP_A|N5XO;5wk;l-k1o`ET+THUBUN9xWJ$jVHP zCWC!*+~ky!BlLVry*E3gZM~bn?CGL`Y9CoA^ya2QGu84ry?i5-u!s* z`m9MzQdRzu6vFuv@jXj|3eS`pMsq4&q*E;OG9%$_ zc`(6TJkDEX0u7ymQgLt&(mpQ*v9=W)QD_ZPM~;Ho1YT%%z?rSJ%G2if==Q>-JXKW< z`3l|!9Zs@CEx!`yt1F6K7HK2ttPw`tinuM6ch@PwiTE5c{C1vNw9bp&V8eR`uOBOD zi8Hl-2`#KkLEWkS=Tdo0X(sebY6W4RQ)q>D2gF*;dkrx8_$jg1+YUq=e8aWv}G=<&@wo-Q6|+p|E| z3lUjmkDOtLk*2hm!Vj8QOy{edC$Y1cUsRNom18r`I_=!MCP2+_X>H1`1Rf5Wx>U{R z+8X%14>JquZ+LRg{x5wr`k^aDc9KnR=E2K)?`f0H_!h`B%Nuant~lrAGi9>yYXELnmB7}3&YYInOfrck^^#*W&P8~_#wIQPG~5-rE*^h1qhaY%<)>-}dZOM?9 z(VetY=RG-lul8h`(_|_U_y*uD!1G81WH)D z$9RHt>>T{P#`i{YL^ATuNyuLa-08MUn>8gG%LJZ31gygdu|3JSCL`7u9FxW4;E33B z#6E!m&p|V>UGHZlFnnuqbx6ppEsM2(5Xv%AMA9x3&^tI+NB#Qau0oA8g3e@bY36>z)SnkGW$M{E5-pJ^uDo5QLRb zK}{LFvry`yI27diOE|8EYl}f%&|&{SJEq=UF*kb~PXUYg4>7q}jX6Yi$UX}>Mt-x# zQTv=7?8sj_X?pT}drt_)MH?~+tvAYI$rN2tUC}7@yDQVzdh+*R(7WrrIJP2;6~u(+ zcO^vgC9g7u!@Hadvv)<@p8E+wjG%O0n;YxpRUq z|G}0bFB7t-qo?rg~4GSmKo0l zXZ39v=PKnY%|9a6LPR;5UPuAvy`zUH|Hqw&tf|6dT3m;`K)m%x>_%mNX#_3rrKkV0-LIi8@;u7 zYw;D=IZ2N8`1RMQv9VfcPjM zey8eD)9Dl4ieJgOrD=vjYMIvrKjp7UNE+^&=pp~$#0A&f!2NYt% zS?=MIH``g_1z-ueKOl=qWI7AQ!OtH8u;GTd_My)I9)l!s^94noq~;+Mrhge)^lvk} z!5fwme2iorPpyQ3c!fE%gVq_jpOb|F0_K-Hk4hflP<=% zqb}|A?4pY`Hl^5@fk~F@{{q;6mC1(AuEm)F^B6+Wb{^xGY7YS*F!0{5IP@&Bb>JqB zq24BVJvfE~!kil=suQly_ZHp0eLH*Z#Dn)2l9mhsmNuxl^LVXsvUZJ%T0slva=4Xj z`N`0?_w4BT1>#3ce{Ng|&MI@U8yc`&HnWLCQ&33R(G!AO_8q}*Eke zNcm@)_y0r6|Ba9WJGwabJc(#|xp3Gdjrn|&qdHZLZ4H+?8$0h+HV)BW{=9UFCX>b_xi(Bzdjs*o!6n_ELhOm=hY1Nu-Pqb~UxwC&2R!;n z3d~B?5sv3kD;U;Rz)saYrUO9KpWvaS1xPN%Pwoc%Y#V_4Mc#2Woz0#7t~1TWWfyJN z+hX#T-Ujp^!$F`kzx+;X=JuY)f@brjkJLuzg~Rdtu$KonwV4%|U=W;{^uJvE3##A3 zgyOQh$;@nNAfK`37dgGp-lntmfl<$c)bQ_q{pJZFX{><2Fa!}wA%RRl`>i8P#3&Dk zK-g>lxw}M`67iquSl!MH={j2*)mOT@ld1NAqEOaC0r!uB(*n-{m&^D%qW-=sb#tlr zBXye>9p7s z1vvnNqxo*_&|$H@TFX2V=*|6Um8<8anm0D=`-Ierj^1KHH-mU*ZYvgYSKr)?N=~Kn34huz5J8Hl5;bsg$ZoqHA_jT;V1ibx!kNY|9RQ4lL2a9KDC`BjvGD3?0T39 zB&m$;J;b+ZV3#9iAcnajH`q$v^LUGsU%%xu4 ziQN*%!Lq6b*H24|sg3H=Q>6#n&O;){Ayny!`7FEnKp_^r`6sqDf&$F{A2F}(^Ot{; z15@umkv3ZjrFR#I{_Ec@a~f0OUZCg*YyzBgOg1h<~*n)drjavDsfq^y(v z2mLP5dR}nTeD75C_Hz)01|}L{rT94O=4z6+(5oi+^=C!_IK?k&h#HJ>)M+{~FvRHE z`WG(q-)_n|R6;_E_s-sSATj8QR>OmSvYsJ+GEjY1sc8_RBcj&)%q784JaDTI8re@_ zbZ;$XK%_^(2;YChE6Uh7kQ?W`m6<#%O!Oi?f7QRWa_=!#0W;N)_mO-$@&Fin>EKA^mhtv7$O*9a z`P~Z3kyv3lKY>ZILq-aupca`qI{+)A=5uE{^_HZlnl4%IH1#;a4G|1*M5Se&p6^Ue*e&34Oe_A@_D?NT(yD0Plem z;Jd`?ipo?cps@n2?C52O$%Cuh-^fKEX{y?)*kXW!`_1@=_PnCtw%UEPbMg4YM*295 z+dkiUH&vBeOW@AGa9iGs;&DPJ+_+plfKE2roSF?$C58MR+uwfxZ@Gk?+Hn)c*iOma z>F0D?+>p8!17xvTzsIizQM3SF7P7CqyNAWujC3}gY!}IDQJa7FyZ>D{bu6|*4pt@{d{@bL?AN(Y5~;PU%12H(*yq^vP)#} z&i$6P_l8ES1?N_-t*3SbNsNX-vD1P5Heg!Id}skuSZk8^?6e)cbFAQ=*8K7LteVrh zn?caWhha2+drjXYx_c-P*$P7)Tqt zD!se$3OK&^fU=yg?^E?d7s+NQDEC z_%TzIf52L;z#!g98Uzj-!33UyWhlV+db4_}YUn!GY*Vlcx}r_rSz@|he`BG|v&*p# zd~-ME&ZizdBRph6Q=S9V7!gnE&*^c6_4I#I$||aGY@KicGeXiqUTQEGcnY0uO+bI?Whl|*#}V6MJ6noV!;hnj>s#v?rXSx($v(L_ zR%_0zGmun|pJ=>U#-Rp7^AL!2qslqfvOz5&J;9T0g6NwseLn!1&-cV-m@cGeV&GA%^8p5>x<<&1n)dJCb9)Mr&ryt+ErF_H zhP~;TwUA&>m4!p__w8u2o^CEvLQ8(%U%1_vjzEIGUw3F4EBjb(cZU_bK(dYHm zdq5FN3C;uPNBR7r&b2x$k;7u(@A@xCzeE33x%}ub7Y`niM3O{b)VK^K)b_%eOE2cG z+YNVq*rDQm*BH?E$44bnfuseJ7qyg6VlUCQ>zC6U%2tpO7>9vu2c*|Ecq%?JY~rx- z)H;Yu4!-RKDFTmou@EM&LR2yYwGRxAe2Gv*>~$G>5~B>QN5ocN+d;Wk8-dQyBLDiq z!IC6&iRy=m(g&!l#m0@V`azTTr| zL89`fv=Dcs{)Ui4`QcIdZ!jHGE=ynM%CoI=)sSo7T;5Cdeq&g-ZrMxx$1T!sTYRjWxyAMAt_iQ)oB|J>)+Zw-k{Q{cgPcG@Lku5($D;VZ6cjNhRAc_&2>G#=L zx5K|3R{DWnGr_2`oTdpDR`?PTg%M%gi`ZodI$_#XKvZ z<}UqJh+@%@zG`z@zNu}R!=wqw4yCr(PJ$%gch;mI-3P+uKjxRHKn`k1EbX2^1LuG> z=12<;c51m~gXgd1`1GfMtQCCa{qRl0n8$zDxzFHjwdkT_xEU#gmtgr))j^*a2BV%Q z;8UDm?IVYQ($886eiVhiB(?ARs1OFl`=1>DQ${)-E*PfF}v0yX7z9FAhxaWpDMO(EdC?HcM0M@Wb{pY z{-;p#*hzX|q;il*S7fDi{s~9BanT9{!uaNBVq+MXnN}MQ-I+=W5~RVGZu3N!>s?LG z%#?pmW(@l^TJRd*%<5Z#tRe|L^ynWtZK0V|iSbd^QN-SL`KE;aQ48?Hd$DxqtVfaC zw%%Gvy*Y^A-I2-SRR5SF@KCy*nhfJ0CB7m4=7-DGv*K!jMWCuYBpY`CpAkM8DsvjJ2K1s2nF}0YM&X*D}bqON+K=+0+F>zXSw>{~J<#3&T=2D#vz0S6WXd`=P+G=|Rf!oYpyZzm* zMXPr$q)Hwe%n{Ctgv9bpDW=%h%8!ei)_{>-{btovKB9AYDwP}rP|!?~y~8hAoazLY zlpd9(T<eAjxbda{k@t$>=L5W}bQC?084v3-*9tY&>(Lgk~y z+5zp`)$XTj(Yt%29zlmZ@4(x?%L=at77K2M-o77m@GtQ|t#;frJy=Vl>p>R}-&0f^ z6zzZIPh?IA{r&@g3L>D=k}Ld&ov7E7)Wmf1+&4D6`t(v7a4OT4lDpoLVeDg=+F~ejx7SWgXG&&hGr&5t7Q%ujEYfWApmp`@ zT1@|tl?DPB>@zj>%iFTpJ)b}#5+v7c3GYGLRPfUUeZ7U>E<3lhUZz=Uk6%2Ym%{ax zZ1N(_m`5!$T=r}pM?#XgWt{)CpLf9h)IIaZ_{Rxt!7a3*(9=re=K!ej1d@d|m&np8r3)UPl+l zXtRH+4Wahz_1KbQ4A1?2@TI499koV%VyL+;hJ)wG7b}y$8sc6j`{5+bSxkk<(D{L=?-K74f0Yw9byO@=ixZ-U&1Oh} zoD5S$p_4;M*~RkB6o|15U1ob_;;_lF>K(o4o`E-DG!uJR`Tu}I@-C3}y_y<(9RgtM zA{l+JYu<$I3kHeq<(wNRqOoa}>kYjbQ&`11jw|VQ+r(91t|lbXIMS`*y8x|5N4k2- zZz7x$Z))+3JFJkez|&<|H*_2B9Uz`N^7K*&qAkRr*UK~kQ)2QOit#Bc!Q`ik5pF={ zE2JGGRs9^;<%`50!Q0+y=23;0eZuk}QuW6``4r?2J!}xVyJ~@mu!_q+gq!1`hM1JPyLQwzy~wXEJXMi3 zKT}H4{sSpcjY<~8l9LG9y75D(9?=3n&zv5W;tq$qru%~S?AGahH^u>iU*MJ39KosK zQ1uLEs)b0^Oc}Yz0=R$?2LYtdioc+eMTLcRk&#?17v1Wc#lu%AAW5*h=Oq3+C_>R$ z&W^fmEBvn8wsAfbNlQ9Cu2-*Ml<13)Nji2Ey7P>AzB#v`V44P<@5` z6tdadQY;SUq5SApW8&LyckCF3oT4DeqtyLGf3D4<3&>ze7=dr|4Oy(i7jGGZQYe-z{WO~2gkHT;8fW6>ifcUu4hH$DiB1U;7S<5o8zg6ggu5~~D+%?~tmBXYA%A%n1x~KIE&nT7v zn}Ctn8#VtT$U~O&Uurp0XYP2H?D>>Nn}$R-+&K;d{V8t@WzqvfokD&7%-KwQ;#M}R za`s4Yw)t2Zg>li?*tlztwzgw}cUMY$j55Rqk!h^*8sYg~wIHZ80PrZ5^0x6HT6<}G z|Fa-yTUUz3)C$tAPsFei93C~DR*KoetXe>k1X9b}tSosYWPgl#q5E;O(PZ&ZC%m)K zO-x?b)N7|eNU8y~RLyx12zWh^DrB}h{qcjWEI&6Jn{pvoLWhS{=%eSnb42GOz^%e^ z99F_`suJh;#d9sO=*iA;^i*%$`cF&0q>tKyDog!_-0J_hA#=YXhxhDUpl9YeZ=O~) zDBX00YrZL;5j2;m)d#z1ZXQf=+GDDxf+i>hBr)PuBi%l~YD%}9rKRq8mb+nfZmM$E zdTFcDP8rK$yDz3~mR^7~>&iqHj9R&d>n9xnIjo@AEz`&sGKliVrUcyAj$5Ph2?i~}}1Xb(~5%$3H^JV__6q!dc zM>iPZr=0&BdUA4weZ!Y8oF0PDA&Mf_+Sj&}izgvrF$rNWcKxmt;CgatOa6S3LZWCY9Der538Lzq?*L=jNi2vA` z3v8!7f}6`O>Rr3*1|GJ;5a=Ph1ltO5ERR^HdN5-f(s*|sd} z>Q`S5&{d;iYSyh@2X!LXuE6C!n=kixRm;i3}{Yj|8F}k^j>A)hW=YCpnflQGy4a+~c?xL%W!_ zk?M4K!@**9@RTTD(Ea=OCoQny_^dT%qa!gP-}e!|>S71p$8Xy_?jsCNaWkoHl=GJc zQ#fqG_kXoYPiZ|-0VP~b60X9{a5gq4oc0Q(5)0c7Uv+|!zp@!sSSObA{DC1^A|ozJ z|B{()yQ|Kxg`H}560q~)1dFMDf<4VF(=Ms}zR!lYs#?MjZmGRjxAB*CR^<{UbW}%t zXJ_4~Y*lY7NVuPH!rj0~TY_LFzh+5Zo&kywFV>7~@_sB5;BM3CMw#>|FjpZoBz8vlaLudZbgoTM`Z|g2h{K*n?5X1%md3K3~-AIXY@#g zLul$X!6mcusb=LbvvHN538HfdNXNMA zL`i4t3Qz4%`^```a7L$BAz?w+XK7|48`oWgN-CjKdio^c3m()kWWe)-$10|8$QJL8 zIQGHPyPY}FztHUSz0@*+KFN%~?bgNcpZL8Zye z&k2l1RcaBaKyBlu3MsjH;TM|oau{?PzM?-4-96;MaGfKxcdg}?{e<^ykWpZk2h#ni z{ZvLe_uRY|nFlw&mijfAC;vKKYwCfqzK4`qh8KQmL}{3F%NI49pRFyZm@m0UpB{)$ z^JCI5W$rlCP?%Ca2VFQl?pa#ja~M6w<#X#pGlt$CMWHxSB@-1`p2`c^z zd0MNDboA^2t~FeZ_t{kDgo=$3bmi>D4M9$%p_qunojDhVg;1{;wj$&)mk^Dv$6J+p z-GebiX9N-%(3~BYjQif|+-r_s%>Qy?rWM<_z$-=2ZwYxpbK0Sg&$=r=IF32HwFO z#v#yi?FmQ>a)?P8H^99@vKQ|ZAo+JNJsjh_&*ked7h0I<<;PscI_S^4%_gBk#K8`o zjaM#)T6XOsMtG51@eeX3$c?nXjcLGacAu-MTNle2WWDEVGP2{W zsc?JUbpT2tQi6wkjszxomHNNiEz)|iaLF{qQf|JVnt{AMke@S57HkS#ge`sLRN1ij ztg)a1tA`G>^#~ea(pB&a`siI@T5Z%hE3a0PLo|I%o7ul$>$Mb4eWSw2Nub{?8I3RL zGtw4T&AGQRwA~~zurmx1P-u;f`%{0+U+HP&;5In;pWLI`>-dg5D z*W<>_fMNR_&Dp`5NfCkHhq}!1H3hNhiGi1O#zHD211LEy4I(>zpOr@@!$i(sR}Aq~ zo4VsFAp@D%_iJpdU=`k$v2X{YNVkvqu01&&yV?=F@99;e zFCsjHe>o?3{*Oet(5H!4bpFX5fWu5ZA=I|x%;I ze%E*GnVzc;A;h_WHgdst%%xNgnYaYacZ&{5lC%~+<2nGv2!zaMd|OE_?Tz>U~#n=QPX zq5JO5g!s#31P5V)lKDn${Ir25PbDBQbdjk0Y-oTlOaN_pjqZT5$%I zn%}+Yd;?RL$OQ|ps>jP-(4>EilFIO%I9|7jpQJz*swiGhjU%|)Z2siPrT(-4fG$wR z!}s#e`&8A$a#ddATZ;|RvPdEa^M#O*kU~ueJhWf-NgHD3@?%4_(!?(P+-B0`5#|y( zCr6O1B0eHc2&qXxhT#3Ot*dWN&YeYQ^J~&8C8e+TVA;;ht)cvU?k}?k5pFN;mYMvV z{@m(|&b)D7+QY;9E1omrNCDatYqMw(3Ug2Bk81+WXfg%8gXc1>;#Aq*lT4`82!LW5 z!@uq;_Y>Bf9<1}~$*V(12AnGB!;`5xw-)V1hDp3vE8p+|gyfteubc<@xmKSA!vnAn zhr^5+PpTZ@vdp>w!wO zTjSk}&GVdwU!&1m=IHu&mUi6IDtWV4CO-K5F_?UfK3G&Yl^Hs_BJ-KV4z=HY#x);B z+1t0yQLD|gun38ovWqjX7j2TS7rd&b8y}e+nQdTteet6a(ODrhVcA{;c5XE1IeIC0 z4U>Bi(`8;ipa!T9u6mQwu5I-O{-1Vq33XVnJuhhl$Bn@2NQW+~KTYFJUIT<}koVc1 zW*`KWOoQq=pBj?}yWBb0LuBI#4;z~*`W%93I^kVEq^#)T@87);N6vz|5q~GpHdLRG zm3GeHS{2^Ig$DP?0g>Wjge6sDR=kpOPCLFivw?OAf8%z z*6-oOnu`W3?R%nfLFDN9A6EO(h5d8-6wUJOlITYZcJD27#=U7bo?liL4&B-Rbn{>d zUemM`z8+REH%V5RBg5{^mCTcoRv{Jqq_y-s27?4wMrB}2cB5xoElVSueWtJ_6#kZG zYR8j_5Kq7A;-aC=sv+v|Wxq72v+8sr?}JJEXIeB9=7)N4!tF5x+5Dxt958+g_(RgHsaQ|v9lXibmULDxT* zBUq1^zJ_mRl;GtDVH@Tn$o;-UVN)G}#ExuR^ogTLa&A%_QG`BLXE*m^+Z?vfnvdu+ zgb$<aMrO`mD(HE7otk*8ntJh+crl(TvZldk-4v2Ohjb+bxe-N z<_q16zEk8J- zB#d70lRz*RI~?ZV%)X$VN!u-h0lx&y0({fs#*_WxX_Rt4z_Gx$BlqLa8^LqmlQ@o! zzX&PWs((aGh4wMD(WlQly6mm6bb}^+N&eV?MlntBeo_%Yxw{{|>6^MGTERW^I)|a1 ze}t?U*h*=1R8w%);+>->P61q`G_kv|+`?DybLcpHX@Do2QxF+Qew>qCp0TZt#9>jM z0Z!@C`CdX7|xM731^)cHpzJ zl|vg$t~xWC17d2jvP#p=H!KCMJrt+cwX&@$<9+(+PhMmpG=a`MceL^l7sML%*83VE z=+DYh-+k)B91f^B*uPL0;<;exR(xyHSZ{$xst-%;#cv-oRw$DG6MaMB)sTWlrKPSat4A-zA)m;<53whdY z1)MgG8c_i8F^2LPcTeKt2VwM9n8&NqeHi&`tJDDAx!cRlCvOEVt|V&bbK^8nL$Dni zJ>+w}VB6UvziSD5^UX~hn*SuYYRyE2*J#qcUZ<>Q&2Q0wja43JpE~^-WEtF-!dWRD zK_h&#fX=G-o(6&dI5BGM&jDFgJ-B!*Mi}LJX9JPN$HzumeO5mn1fCEY^D5p5Jzl%= zxW;x!zpvHU2_2@NLxH;-m__0;eJ5xxkEw;=A^R81^d9UP7aHDGVe9ydBxDE&%x8|g zPZe6Ak9Ny6Xx>cAf0+jLY0~wZHNlxZdXm@uv&_2P?mxk(WN0zEiyp94nahX%wvqt- zw+4&Eg{&HEtoEgWs^8pZ2DRKuZK%O6l}=XDpKPJomx*3wm3Zx*yAtd~5(0+8@(&1s zJLvrLkgWDs*K@mB3S73Gd}1KoLMH-Rj6;6_6VcG3?>|XvdjRL_F9KgA+OLRA?#Z>S z%XsDn4_iGBQ8__&@<>?G0-4BOX+b#q!87yK$0`9?h>H1G{PUZak`84g%f3W>AxGoNo ztah>_t9!S=2g!$$#0bvF2}q<9F+%W;*iTlltDR_JydSzwh zO&)?JHdWOqVa;9i<}#siGa0x5B{VH0SvX!cu{=&D>c=ar6;h`|3R6qccb0pJTnV!X zD9nP4t*AIa!~(njT+HQlRRBi6ejV%PKO9nETKcxzvA=AC;TfQhZR1Re6}v;Z*Y zK&qOnArZU#NMPN!JB_D?8)j7g!Je)rBj5?2ScA}$Shho)spdH%-yAXMRx)zKg}XoR zHyC0guX>y0cGk_O`k*J`MW?&^>EH`+w%;wo@dNX5%3;>ufBSRkme*k)*oL#3+gTTt zyZ}Z*;V^A6_c*kE^TNyPS4ZmXvRSP+wneDRYSIt$m=7)C!q20w_vBzA%yCqIpE?Y)k$iR zr8&B2Pl(7MMfK3IfQ-fD$&)9QDZU318J~uA`f?ESEIvLfIb$KS zq%0Cc%U|Q%pXfLBj$7ZI9357o%A;3o4#z(3DZwH)sp>|vR;Tg=Ftd!8wPS5iHSbA$ z6VOFDw4H=1w$E08Bid%#QC?nOPj^}P$>_19&JHl z%CIzvzH)3aSNWv@cIpSE(ogB<^Aep^ki`QMR36woCkkF#Vdp|9saw*bR2>PZSAlure|-6A!@!@al3%sHN9rrrr&_U5(bB+4|>z9Zy**PB0MA$ zmd1kL7b;sg!vZJm&xeLELLC+Hs()t1v1gpXIs_2jTp!RwRe~2y(X~+fsz^> z!PCf_5luoULHdA}!~H;d*+d_lUJXC(-HRGs3?IlJyQMW%gg@F6)$)$Uz+4*bR`YsKyX5I738eW^Hl3U-TMHuhWW)DDWfVO`MpFpxrMyG zZt1$xnpqz1Yh1m=$tA%5(^LRw@0t3>GdJ0t_nGL`)0d`HR-EQs?lVm7J&Y`@l|keZ zL`$}T>ZwCpDJ21K3Oe7oB{FQ8-{hy-;DIhzj`UCC7$x6b=!4b#Ep(tu5~|Dv>=Og0 z?HP>$I6bDbB!TldoGGAPG8B78C1edN4(9^x5qP>Yn;@5s-OW3vm+}6tNjb=2tws)) zI@fRet#|r<_tGOuS$7IvudQHL%a(V;Xc&1@;4~F3q}WShbmTLKM`YcYVirgm4nfED zMW|7B`*5kuIo^09_$?In`vDQGi3CBAJ}Jdj4u(@u62C{}rYiL^9E>H%WI5_&z51ba zfDj}W%J(x|*m?SC+^K%I=sZD7WYMqy6;t6ebNq>ao-tm0@FJSv=tzLOXR3^@9WAW0 z>J1z%W}ctY7MbBw%B#;)HD<;Pbrjl$BczClR9HK9_NakkLP~|-X9u6Y3QNiM86<)k z5$@3qB#vtBP>q+?qu58uNWO70Ph_|1q3l47XZLUGBtd%Qt<^!RkM)B=;IFgIbKa1Yz_yX3->pRP$)DT*17+%|pcp=~FTH z5bw!D3kzu?BnW9RJZZ)6UUA5-RFw{a-4EphJ{)o5bSXwXnKtR08&_2hQ zs=?2&h1?QG_ryzO2?+u3GRkfQg^T*G;MlY0O?jF0@RH%7f~b7iTMGlb=!3x`Ae3Hb zh0^40%Qmc@T9xxNaMZ-An6FBa$-U+%CSKpHJ$?Q8*>2&aEsK`h232k^5(!e>i~fPb zW`Qq>j}|^OyH=Rvpde9~Pn^otF+fH5w?6Nk^={~MMU*zt>&zL|#Pb!^GM2ejO$m60 zZ~S>9+{_;}QxB28FU^bm?K|ESE>s}V3BGvAY|b@-Ar*6G5(L`B6aNLxE!v9kz4o{wpqJQnY}}3B=ITZYFOzT~h9BPmE83J8I8(w6 zu&0M3PR&Xh_5u8GhS^*}O7M0F4qO%mfN5kLfMfuZcM--8{(OuG29=?&weG9oonuin<)qKa&D8+u{T0O_k7EqK% zG$+3~+{_%kby&MlwS-zY=acE<)9V@1ll#p39#a-acqNmT46ELYsD;1w8{BG;u#m(G zOFYbJ{wd%Gzj2U?r}NXilnm_r>p?2(ZJhia^>9y+haP`ukWvojP?F0cKmv3VtCLRV_4zq`NV zNQYA6a*aRbELmTjO*yCeY3I)v5kL;CM5R7mZSBY`Z!*`Kcq!9&e21QsI671CALABt z%KM$S&$O9wnPOn$fLgZrJe2)M>+?-(jF~iQ-_9(()Mz&0+5MWf1T3 z3z`A4V}HpJ zdw>uT7FMLO`RF?ttJ1cxf~nhewY7bvezq@r<(6A(=AcaVjZ%n@8}-{d%dw$7#ojs- zIDI8^`6`-!98yV80DcfRzlR!>_Y0| z<_q#`5_cALq7V3lzyj@@*Y+JK@FTy9s?3ePwLBZOE!%}2N2_WI!xBheGAj!5P&gFc zNn@3I&RVZU#i#Dpzu2DX?;n3R!r9rq6ZUx5n*JOemHTcf%x$z;*wkyiDXopD9H(|`w~masseTM;$}f`jUF73`CNcgV+JTGS!LG*#Rnhb*FIqpu zd=0|npvcSIrt&M7uQl_dMnq$C7rt=jZM-7KDrO`Mef1v0Y}PiY{v&iQmJJUl95%}M zId#8=YvWc>T{8ZqZlm`8j&D}|X<-otg+mc=F%`mF+{S1a%TH{Kuk@YRK3%<8X}4UqLsRS(Tx~q_1`~h#_z`kF6NfZ+fMpE%aGamuvZv5J4-ZxB zvZQbmv-2fi!OC^Vd@9brOttm=LwKjXd<}k8#YxcOmdmbDjhfgAEK#mMZNb)Oq)Dg& zeVsB$9Z{;NiSzuhuj{26wFu`46@v7ELJ(yia+6X?S|+2ER);J7fCY-`q?Tz|$;o)) zCZ*o5;VM%{sDDJ47JBR&b4PN0BhP%0%5(`^!&9VGHXin=hJLsi&OA11i>kRZJeA#W z;IS#Aid)vJ1=CN?!I9Br^9#-w>+a>WsP$O-k4t6CkNu%@z{9yi>NL6GB3rp>4XqCiE8WDzH|F=tx3OH2^???SMo@)#s_sUQ~!rjzBu41tr z_f21g*mjhA^h>~OBAr@+)99gK0&B`|a3ah6ugc>X>;?XF)>HXCfxS7;TY)4JR7ZMk zS%U7aSbVnBh6(~X8v(Tk&eU5UqIrh9SF;ny?Wy-^(O>FB=Y{(m`PAowK%#@9*4v}_ z`6dLdDIyQE@(}AAcCe=_c~}JV?1jdtd0R=NAM#?J zY46H<<$bJM71W5ysGhb0s%m~ZA-VrJMEYNCMiaPs4zXm_TZ^hYWQ({UFyplK@_dtS zUYy1R3$0xG$U!(jfT272nvwmWOy=_wqY&EzRX3K51Q1Juf{x+kmwlGP1;=K!Ex$#K ztnVN*1;R5tP&xPeuJzw~h9kIUs?6v4@zT@U z_rFDqzwfxPF4ee%JeZK#Q=pK$&!WP(jM%wod!j0Q_>IaDe1yBqOQJn?h!nj0BQYj= ziBCbTWht|-VmNcaCx+I`eY9h0J}w|USLIwSf&}Fg(YfcWz}ET9B&XjaEm*uS$nl<} zzMe_!i~wfn^Tn6X#@s#Hj}de1d56PbWw>yU|6u8$+ic>S z82sCA{jUu$N4^j1`1+MTa);E3U?v+|g=&Pew4TYd@rA#C59i>jvhf(r6jY+iKG`JF zZOu=bWe-Sx?)!TP(#Pm`nY4oU`JGb%!a-}V1sYWOe-+%fVMr2(|tT-D7YeJ6TOuAL6-$%msgP&Pk8@8?6`HnIq*}`BI*CgP@V1p@f$@uYGj|Vl5;-9&n3uNqWW&}NC`nr^GTGw!HD~B;>Wf3jm)B7eu zTBtTXHl5$k0G3~m<2rJru%oEpa_g3YOK-)g*gLX4V!U9x*QyX#=o{K~x7-2zfV*i7L$+%92!MNhp+M zM5su#kZd7a_GE^}%%qY~X|rcfC`;DCU}UZA*-c|Y$TG$@gBdgPoHMSjyRQ5Dy}#ez z^Sqwt`K#A`U-*1J=Xo5*c^vQe`#4TqE`D$@JeKiEMByIvPHi&YPq%WHDHf00wxu)t z9o1-`!fV>Qnwi-GtGULWB3h}t*7gN^I3{9{N}VqoA=2lr4c+gnmrT|%-qrTHElJPT z9cBR;2&Jwst7(YNN2aCXLT=&8Fw2`xE9lgOQ^9e33qsqSicHLS=R++3zF+mu@57~} z1Zuhie=73u{(mU**6+~nO4^2MiF3dYR>JWm9`W#D!6`6>Y7(X{*06ted3yv1->1f_bEmW) z^+I`TYib_|TnH?+%(q`TWnewp7;H~CDXO0SVIl))Xz(*MPT$aEz;GstVtsq9VXdP| z%8vc)DPjAvE*xV|37d>(4_oNZ^7=RC1OdpBLd5MK&kdZ#jlQwYbG0_0>`hY%!*qvl zlNOMv&80%c@sIIYgc;og8R4lpb)pvcH==|=Y-SQFc}Se^N81H++FgX-=Gz0NVp=8B z5f8XkR+c}b%u{+nB)7#5YED_S#qp15=0Gtt-MwLr9S=sbSCRz=YR1F8O1_%3#L9Z`C?+B&!p$qL$-mpEY%6^2sn|rwl%Aumy+BY%OmRS@iM#cw{IwrDf=peHCw|2K6Rmb;XlBDf!I< zP47E@{yh1>EMN|0CYKJ}xfA%DuPZGrO|wcVI{tOnrmd}AQcl0)0_%cg1Wfc0sQ42q zW|i>sWSgA$3r|5h4AP-ykAb3;FS&c;O2`fwGwrt%AL>+}sfJwtBwpzt*P~V_C~H=% z@^i*Rw9*&WMxFT=bS?g6RJSahKQz-eZ1aY-7_RNR`EgPCN&M@Xua`D!Usv0lnaTjX zmP*Wu?|2S`yv-E_#*|Ra{Lm$B@v1XVA&k=i)2HXC4#6xiq1If$%N}gvEEm{(;%k#A z(eH+qVtz7NzaeGD8G@0Im%IaB5LrCt>&J+&P~JCb?#>@wGcl>dZ;Cb7mFxYJAflcE z{E%ihc=_@x@3*@p-~zp-9)%$2+P2|tCF(aN35uxGpEI<`yfa@7Z55~bUS7MMXsn&g z#aY9z9Y{)1Z|>+&g!S1OlswhP5HuIt8^oSLytlA5GTK>d*i-D9sKNH%u%~rFlR*AJ znog5^sx}~RReDt5`P+=*w5#t)0oQ<*mbEU+BlYIF8qLIk;=+u^nFoWc0TJ;^r~U&? z#-@hV)nTk6$**eE@Zn$_d8w7dyFHnB|L=G@3$bN z^Y10tIj)y>YRK^>TSfuxC1gba)_Gnvl9FM*2{7_*ET#1~@So;wtY1w1*j%xguHdKO z|BCw#@NS{^Bm%%$y56-mf@_G>mY}*GE%WH{! z`L#qmrV)G&?{mUKK3SG&i*8w#gf=uZye7nAT{?z)u41Hd2_&prf$#X!su#K}a&KeI zA6gfdS>QaQkQtyK&FgBE^(fH`1$VL!xO)*Q&-hL8z-Hy+E>xD1+WME6bEz+dQr39b z>Tg@7acEU_Rn&N^b${(8CiHTcIGc`W{gaM(yMtW$lMaOPes3j_T&O<31c-6{t;bJ< z{4QvcyY44Hz~hh^4ArW+C$Ut|(xHdD*Q<5y<{pOy(r&S2LC#MD0)D21eB-fQu4yi^ zrY&sy38aPYX{091Ao)n|Obe$M4b*8j-Xoi6z^d7FV>zH`L(`MKC_WUC%#cut-gc0r1z_)#Umns4lzAZ1WmRUde%Cez7#%xZ6 zm{GYQg$?3xYy1i8SMDS)r24qz3z6s@N>T~8zdqg^54qXEQv{@}@I%*VwsXQso|+`+ zYAb+HV?`ZbLM9uSOR5z^+mt!{=Y&NR_5(FbiHJ`fzbH-|2*4fC(2zUfDpfPGc@;|P zCDM;(KK3w4jT&Q$7W{28n+G$JkksrFi~R?g`}RGW9{PAH@vJH~awR?C0%|VY_V#o< zq$JCzs}0gTF`jR{>tP=-jbAXy$%)yzh3tnsUf=fFcB z!fv8U^&e*@3B9eozn%Hy?r=ZAe+El^(g6IxrL>+2)b?f#%{b7alVTxIqx;so#nCXd}+Vwmka9{?rKE-Q=OK@mIqaNyq2XS}MZ>>Yxoq)erZ5deWCq&P}cQvK#)df*{BVq>Wy0 z2@}WJaK|*rT>v)rW#gY#au}2wM`$c9@YC@scz!QQZOM5?w(;H0h?y;~PKu5cK#A9u z-kDc^8Euoe0x1g!>Rtf_)PJqge@M;;@FXhytzg$;1-akr%Lmn z5`UjpHLx<`c9O0^0NO6+hgb!)dw;<+3I8;_{&7Jla;`F%8qdFVRa9!;{bziGYKZ9sCD^f*G9X#E;}iM z)|;E*LHfWeAeK{J_V(Y-8;PuTycxDw)*CAKUo}*#)$E$uNM7agkS9(ocHlgn3*kYu zjUYZZk1ja>gn>}tV?hI*1+);3dKGHK602XKY2dX`Nv(mqJg?3?fU2i!b)yO-Bfe1a zEf_rVI zwEc3=CGU`ZKZ<*wK&9*|K~EpbY$Ff665qN$ z@TBle7dJq+WY-Ty?lx=d@C34_JDitIQGtQk-Jugvo9(eeW+aULOhPxjcrxJLE;i6z zy>Ih}OW11*0b$2Cq$2Bg`5b_&izlM) zu)aQdV$CB?B>y7)Le7KOFRcIfkFP}pR+kBfoDmyCONN5q`8UzX*Pv|iM$#*@ur*%u zUy(!-PzRFI&1+#KYVC1BIe#i!)IW{DrYmdy(Vlnr$-304Cjx6!p2-F93INLeWcD9$ z-AqNmy6jPp&SfvMHW01+Ln6m&>$v`XIpLR-ty%J#V{+^KFFeDiT$wj|9UxtMquao?LgA;uch3)_1U^q z#=4ory-$iOCklhOI}j?9z4dfg65Efp3s_{ceLKK!Y;+jQy=A1FTc;jJP;!Q z6HLQVLf$g{ef_krYlCNSaVFked*Y72%rTR@?O%$+tUw^0a{osc<~tym#8S0VgV(%| zr9A}3`19ibb7KTW$D^aaa}(=0H=S8$in=?SwTFFZZ8A>oPi1JSYH2x9PW{~tnSKfs z$bVNVf7$W%yR6UKob-r1!C^?OQPR8FeK@23bq@Hirj_YXU|8G8W&+?s3 z>SV)4oqqr*68nL}oo-FdT3b-xwFOQ7M+*H9HOJ#3KBN2VK_e}YY|_tRLK@cmJPE`_ zpd}{xTIpFMZ{C!V1wEg8t|aOtY%wyPadxB6rC7-xo@Pg&UkOgoDZW1nNDsd3ECKd% zzcZ)~r%v(`uN*X!hqLu?|4Te@muFsn^jCkO#iwQeD-m!o0!(CeSFs8_ z!VJ$J?Ebo-MX+z3${G4=ZnS-kZ9Vo5C|f{&z*`p+A_uJPXyc#Yz3>jOXu8AFackHc zLW14K=5YHaYuG{>>DH=oy$#xE==N!P)#D_qo$RmMbDnCn}|iY=f-^N(#{{cPs{YjKZvG-`j%uYFnzf5I#_eSOVeVMbVO zZGHIO{p#)jBSCDKDgf75TWkP~x)LD5>H0o4F_v3)87q`_K*!M8v)&6;eZAZpss zpwA{Vp=%1qzmEl0Lp&ZMm436akxp+}i+JVh4JVH^0CM+m&~x`2-?@T@B*DT9wRg94 zDBgp!fMD5@M}|YUcZJoes&(44w^aBd$G6?vvqoH00Qm<4+vc})l+KFG?Q&~||Hypp z9RJ0v|C;~^5X0axe@Hz+u7wRd1v!2^^0?dW+eKk5dnZ-S1A}i`)3g0;@7e-emT^X? zW|wZ{xA#|CB#T%b*45s8v#C{7UHFyID@}>)hR19CXWjn(rL`qkj0J{F+ZxlvBCInx z*E54TQDRPX)O45EPh-UBh_=?Yt`Qs7$ZEi!bBX-lpwWTweHT!9#yI>#<@pMnZRbS| z(X1Cb2L8uCLj3s_ z74X7O*c7t+rVNXpUKhj(!ot7hcvV~FZR$Xk5vl&;S$b>?}7;bA3lCIfbc^V6V z8T$SvzqR9#4}23KzI`%&}x?>jhFwHQSdZIA3{Yh`R+T|rLDa)R`T$9WQ`nL1A- z$5I{`2Z0<{K^Xxq=48Yt+|0jn{&9kO}`MnJ%LKk^v}{JB1Mk{|r%BwEFn zRtl5wJ%w*2CzrM!6k0P{ZlG*u>C#iyPA@RCjD0RpE_OjEeItdU%D@Z@RsE}@@{nZl zvQA_5?#dqr6;z6UfW()?m=n*tm9+Xp>O7_L=tJB>AN8zSaT;gR(xa@kJ9UPZ_d39@crw4E&v`3(Zq;JE!>Ea1H_QK zOO9?>XxR{x2-^M!Rq}LN5#SHrA|aP->!IVxR`PDn(tK)GzAacswr0meSc*6Y#job1*5sZ$*u1bjd7;5n}6*!W~(Q!)IJ=-4iWYP7)9h14pq zD4`9ki3DwTN=53LQne>MqK{ZRA^pWhmOX1+F;n$nRK~Rg;f5m1!)e?dW_I6b3?M$h zuYTB9f3j`>8ceWqkEHDMO}pg2=@b@sU#+<1oI1uDV~Kolt6ot8Re5yRD6Pje(1aF) zwk^0}F@-Wl5AszxABm1V9G*!Kcx=VVj>0lP@Xs0=9`^1Lt1AlB9&8E?;17A%{4y;o z{6lf2bL-&_veyAiuwIJwamnjBS9`Np&2yGpBDGaNw(x{8dR?6>m^9GtSbrZn9O7mJ zXk}&f`30gJw?A@b5xY8LuW%8`q^TC+eWvRhY-aS2C>b4|gxRZ^W1}KT(IA+T8TnyN z=LN_H*8;VqD$*^QpFs2bc~3!>^Mx=&<1&-5Jv?2CNf1KLHNUG<=U<9P$Hp<9 zW_r~;SQG4S-Le)}Gg&$0ztqvhNY)gKMwXl^5wWjVVU;B~z$psbWbo$A(&F1@WyJ(c z0*TmChy3M}AB1*?HIjbhVFMP%EHFHg)+R#@y@cy44m>q^{v#ruGc^R=gp|N#7%G80 z)$5!Gp+UzZtfkd&R*RilWBs)wx4df6wK=`Du~0DhJJdeVbgMgz4!QxyGn@fKsLabN zr{h&+UB5zX*z5JW_>xAE1O!H1Vf=V0uBp5feEihhH3oq`2O$m$Z}yUCcFaINx5M_x zu8cegOj^xb-W6P6b2#_-{M17zuguAu3=UF(-;=2yx!IOVr5x=Oy9I?2|0@$-m6R$3mnDu&!iYL2y(I3r+Prj13d-PlUMyi&K;D}a^TEg`g~7~=*f04w;U{hNa@ zH0Z3u3mjqtdvWzj-=>4(=e;le2_Zt)y?=OE#JA7F#?P^rRHaG0^#QN@h!3+;WzA|0C7Bh$WmDyxnoc8zdu7c8ZN&b6bIxk`G8e>ew&uN`?-w2%`-IIhU2VUGSr}Mg zP}En#y{_-tNP?y!ZCn)Q4~JZPENL70GCcX|IA38eEpP zJt+&Lrc@2%Bp8I8Mhy!CS7!b<8VQ?6emlF%?fFwEASd)ct@Yd|>-?U{A*8b$^?P!y z)it#|N&ZF$Z>9x`PIsXrw#pD}h+UAem=S!UCgxj-hp+8r+MtZ;XGX)qTm&i*=X-ZS z^$?MBB~Ww493Ng+yFFLf6zlKsT4=wHW_@|BX9ZBzFybot0;c&NuHo-n!$3Uj9(W7A z=b*a2-xi&aIm0B@35X|#7U^|VSXvGwpz$6V&?ik}-gQCzmo@ zq`adv+kX~q`B-2xN6?r}4ZsxZGM1CQkEJ+PUerj<{ON^yGW#K;^Nvhcr)SG^CfcOR zdaJ)84t0G6oO;N1-S7$oLppX?@mzw~zF;m1MjBSuOb+r19gtL{j+8=)80SA79#B9P znWzV8CGC|^xo`a`9ZHD#*`?p6FG$|dwyGOE>m^Y_^;H3(5@xEH_1Izoa`@J8^;$5NB7Z7!b!KAq{{7kH8rxn}Q*!6i zo8dsPA)6|EDJo)5$ViN5C0WYCa-NudFH@ZQ3gUDKb?hQW6@-Y3H3NCe-szSsJF{u# z9sem5kpd1SxRc|^1FbU=YRvO_fXOWBX-LK2nbXai1LNkl7@thw0a_7gw&Sk$<0gLwZvQtsSoRC8M)= zl|0Y;9v-l2+*dYeEZvGqdn;WIQ!mNwtAbW6Y^CMfo9}vq!lkjLZYHI#{mBr< z6^40foRMM6u|Yp0NtE#M*yCm6J(SvDbzbfnjg6bXe%N;SgQQx)K|`g3rU&YIO;<~c z!z;6La^4m84vfFPI5%sq{oSOP((9FDok77ubLKmR$|&)p423Yu#@v^n346;D(ZBPP zGvw$|m+hh_Z-NZYiqd9Tz~dqyi)o%F<`AwGTDuiTtv6kRa;IfD%({BTT=K5jlhtUer<^5Zlm)jX&+f*f3<@BVx`0+Y!w z@jeVsk6%>J(dQdwCFr*mf@Vijpl=*8sBq~%LQ{Xg#ez6x>r!(xqu;Vi3=F}DxJpBYYK9pKw6Md+jpl5wdYsqWK*r~BtIw3p<_x4M` z>YUU})6k%MM%54G!PVp}0f$=NQ7TFmW4ooEkV-nF2D$pYtUtsv$2vw#dgrZLT7G+@ zUQGql2RDn;HQNdExl)&U<{PEIPV3csesZ||fghu7PhbQ(ibs!rR6Zj7V?oTe*fQh; z#UZh};H3^E_giC-ZC3Dzb2L;DlPm2bo`@cOqA$rb@l*#a5c>XpKQ>gRZ4DS4RK!e| zR7RCMx7qVO8DS~|k;ul2@v$rRl>Bd&7B=*N-F1h|4yRUOpLt&iQD35)r{_GeKJYO7 z1SI->@f>U*1E4$~M{g0Ru26TNXpN0gd#1TBc5o#*N#L}*ZCfIJUwMq@ZY{yJ2NJjJ zL~!ZLr1lf~GLjgO<|li5c^1ClKGHCEn&z`DQ)b<4X#<*m55ybYDW)b|CwAg-2PZL# zJ4c@>lskW5?oyJ2E7xeYe{eiD0n5%*D6^fTyX_v3T+PfHL#D^(EZm2i%*|j(Ga7t8 z<2Bp(^Sx-x0poaj570c{aJC7~N?Hv$mBLKO5zOWl(|Em5!u$G=RK`3tOl)9zWE-in z=ma@I!AdR!p%Ct)t{R8Xt?nL)_S#jbbxGZc_Th6x=!c{st0JvJ-ZQUf37;3$^|C6% zswAJ?Nh!8x0D;gaF=HDQRj(2$vQ#SaovQSjc-)WO451Kam@}JQj86EjTejc)hyJxH z;!d*Eq3xsLIbc$yBZmCmNPpk>rF2J{Z-`kxe1KgIuHQWE$q~o9JTvorCSlO<;=viG zJ$2d5a0dt8SOiVKd1XPo*rO#VQ^6rZkR8RRx3eLb)HB^-pRR%S?tpz-E~bPM)cg+z zCkF>j99>m5EXh+vr>ob$SBvxOcwlGIUHP@&l;_YReC@a7EcYtHADI*=oQoln#E*vX zLSCn@R=DgE5*pt4VAmt6j!DcK`i3+?qrO`s@)AId!G8ev-DlBD87 z^pHA(mAs|r`DPT?x~QHTP4Vmd{8=GpWu;6Gi=;u@KOO!@i_dMIpuxdGxtMX$SHH9X zE#uDvt^0B#MyT$SGLs4lc~1D#U8E>AC-xKH?Yo||afwV8CaPZQuo)ffxS4m2wsQ@A zjW+PEyXE52Cj!xO?CCM{^+EZ*6rp+CH`M#qt+_<0dB-M}wB+ayM<+_w_k0Jb$CdnObxL$}bb;!Y zo*TjWe_06IvDR@?4y=#A)lfxOS5k`GB7hhhD1PtcIiufrwd+b376F&1+4P$=Hk_|# zzfe1ZbJHg3Jn!eiOSWS;8O|k{0aoH4tyl5&=Iz_fB5K~!0!twai`r&=gucE{3QE7$ zoc$f3R&X!Bo0vFgR7X0lt4nL9QuRv|U9PT-Rr)- zP1(7*P5}$%1}=;IPadl9m@Hxk+FP<`akQMJA#O%akQ%BypG>}%?+ zMmN1XIhDlF687!u8u$J|itTIP$T=xDiCf$Po5^-&;(l^`-3r(l@Qo++su@w_rKQns zw~cnE>BvaSqQGLxVXf0(rOY=f0`#1Pn*GgfLe+!H!>PD%nqOf_iKqL`joa52>#pR7 z$WfO9i~d@Mzp-$o#V|@YgdoLk@d@8z;pipY1RoP`U>0o)0Op1^+aB)QHeq z!7p4Bze&6l$hr|rA>#UT$667sz&ax%5hX*za3?>%gH0O()CI`zaek7MhyD-+{N(0< z_ClxfzC&DgLNB#fS10#lQF4T>x5<;vE7ww4De%mOAn+Pzy0}jOXCIPE-v$OeJ8)hc zK4+=#X4}2mPB8@ksKt!uO9P|8YF&TMe>H(8dvp{x|3xs0PQO1^7y4O)F!Jr&lV(za zt3k)XHT}V_eF(DFOSS_dwn9~8V@t7HQar#Q)W3qjeHwB0diRw_o#5|B97>B9zgBe( z&N)~)!Rztbg%~@If#k9;RQ_CP)aB&we?VX{V;;Z4!=su;`6v4$E!j?hJj+<9K)}1L z$`@U6rz9q)>b9!jKkJ&vX1fyk{r&RakG?O?`yxJobR~*EsCHX8QKw~|sI(mfc1sD)!eQb1t-tqkZEIa3wbMFYrc*THa+G&9PfvH7 zy~ku|q*fWfI@J{&l4SO&W5E2N^e}5LO?vyf*{dhg1^!{~6Z=tTFB9x>E_SCGYmZ;~ zunfv@?DD1ftgM#-Vd?dH(4Xk7FV`UVzt>Ncqq4Fx_5(@|m%|LTu(0slyPN%@N3TeH z=_=j^UipDL&IOTKocC^R>n8VLSfBYm?*0Ihb;)|;n0x0!TAIVy&kl#NFOi1w-M))I z`4<*DZ52rs3s6I{$yz!8%hNNtf0Wb^J#;C>}V4!}fq@lW%n|~lquG1_X-<^#V7^B5U6>*{Cn4TWRD>0ewyC?(zy#j3fF(Q(7FI?O zSmGTDppg<5&u23e)VyL>slC`;^udQjnjGi^tfY-03*C$}hCf7}5O1|s9ZYpARnG_w zDNt<8F%_?naA1Qo5lL75n0t1p`z`Qn6!8YyamUyqK{uiaF#h;R6^P_Dq zdPRpySzB;0KUeUodjdPU3}}CkJpHNUNlyrfF$3V93l06`-2C^%=$@)~Co+jb;8;ew zw?-wBrlwS-E@(yJtx&1CwL%L-*<#A__xnrg`zPP4qYb%Tc8%tQQSB@fcI=saG5+XzB!=d|bh1r^59`$>f@}%%b#=;)aAF!6WK1fJHGD)~JpxelO8$!HQS@vSp zY>g`P0m3{(U>@%kt}qQ7P`%Rl9{x!J@o0u>UI)1$K$WW)yDhW*I9O+|AVmAM#eI#T z--nKy3h+KPSe})HYP!ITN=tjKE?kT)!I1p2+d`?iUOA6wv?Ux@rDgwydu;ABO!Jp|ieUp}6oOg5kO&ZH3#N!Lo$3fcduCGst->d3W z@3%8_CHH7*m9T?nriJ$Y$*I1F`x|gxSVWs~_u>H5`K)YFpX9f*Nhjk{@go$1laujS zDneRh0DA3--C5i4u~&3e+v;2P5$}0TZFML%@iENsKq8|h%$>npa_g}YaHyPl6|=|9 z96m_OpfktYP=Re4kri3|N?SKtVxP>~#mi46<<+QO)6b#|J-wZiX?w{e^ zx;5rLS42#N7S-DFWGFid9LTA*pX!PNVXu1=N8SOQ5|1LSuKzL++%koX6><3tj#^Dm zOa$JSu^uW15UZGvF(|9jp%j8)P!ChoR*h}qq>I$r}~sr=}7ec zM-Q_$OwYJR@2Dp$?>rf=C_B*aN)zGU*AjT`z~$p3z1c^4OOkHYk6%0Bq8iRWE=9nW zcs;#h5vca@TervV0j-t%cK64M;>&XmjUr(k`6uRz^AN-;>v>1I59Texr@h+?NAE#1 z@-&R-Mf#4z5#i3RTccM7w&)}Gh$n&JgbU2(&h>&mbHNvz=1{8d9sa8U6K(i{on|e8qJvUp1eMR zX146r<)yx%1w7FZIti{SI&sfKi{9NWSXfka^;*d3W>_%jCfaYBYVl<~y4r8}tn{X+ zbfH<;-G}yXh-t@&Ap#24qQ*C7>+4K7w}@NyEkGs4z9JuFN!8q5r>9c}G0xg2V|`bM z*u1D06;o<><(5Bj*+(S=_8;b>B$k?F~e zG?Th5%64I!Y+WdNB=5(J(y&#_16}3R8pRObnr>0*u|Sida6>)U1KE$r{7PqqT%7lk$F@M$0r}3 zg+tS63r`>s_O1_Wgd?_0z6j`ieC0W)o3<>qA6iAYs63p#HrgPrrIC+CQ#Oki39cNV zv|}F`ba-8xzB3}GL;sw#ivG%_+mah!BJSw((_Wtn|3M1mo!Yu0l)mj2uQCxWsPlHl zOKJ~Jt6woXJx4Zdx0E)H(EI_yI59y^ws?GY%(K7Yxf3?XF1`g>h#-{1>%DUYeP&S# z_E?;r??JPSALX(30s>^~go{3eFlzJ>#ETrr*Q)W+h#*$FA6-`j?R}GTiItaL**iTD z^}jj6&?}44QM(EvuE+}xXsrBr*>C&}t3ual026;=!X#HQ?)k)B5C%spw8%v@mk#Vd ziQb=tkFMHxNaf*u1q2V*+H*ciEMtHG{9tH>lZpZ%`p`FrD` zK_Wcwyg!OCL>IE!vfYJ}M6vXvw>oBQu#ay@y%#*j7!qf|ZETI66{#EkH_PP%dOhW&pxlz5Vqi_Te1r(=D%K6(c>K@9`S_tg}h?as)b&tS-GNu zttEUY2oSc(TP>Bj)91*P|1mmTU9x*dT46qq3FoJ2-onybX2iovSA1K9Cc-bKMQSd~ zelf7!$?K>E^#hoQF7qzc{_8bCLw@ARnag z9fgvj?%2poVyFbrp|@~feD#Tit*wIlSESt1LR;PqFK&f5P}d9*7OrsMQ4vS@FLU`T zMWvdsY-d>{sd01HacD99>j2t!yGp^AFBivLz!9#n^YtZ>n2Y*T)vr{dqP-f&vVN|1 zb)`9Zcyxa&zkwhVP2$8MJ-2!qcFDT4`D_V1Cg;?y8EKzAsHZYQ-;| zl}QpiS+^^@-@adr_4v>fixW8c1^y@)T~$g;X8wb{DeaztBr?(%RgqymRzXL})+M zSd-F+`V?H0`bsg^OAEpr(vLke{iR+ku^zhB5J3#&({c~sQzzs5YF|m%eMvW2w@rLm zFnud`CuteuRLf=I=-WYV!@d;^)fwOtQIxn?c!{uLeFyL35EfddHzG$xr?4Hm{e(Z&ss~(3KR3jbr&*Zn6Gau-)=ec3>-bmRvSw^u81-`gEa^W;d zye7s6g$dzh6l0~GYQcPS^+mUSt8Ij8*G6N+*~iDj1(;$Nho@!}jpq06#xxXBG1foc zJNc<#Y~9GOFL(7$cVYeKbSB>JKaXOK(l|JGv7TCy>#=Y4-2NLl{VClX#Fe!&`}fj} z2dD`kD$a&&;t(y@l9G~wW0E+`Se)n%A8dcfEjr$dF(VhcE*lUc5k2({yAn|*smMa4 zn;zQRNYRwgPaA$Cz!~`Q+rWvBP3gyd%9?A<+ry#tWn4*&q*pkMNel)TvN` zo$0-msVh!-0u*V(l;P;Zaq2q=%Uk0|yJ2(ogq^o|cUhkljvuS1Q_cq3(}=l46s9Wt z6DcC!sLQLN^r!ruq$?dq?g+)Vl6IJ4`@K-p=e22Ok>NrFE=KLc*&t8hmwH}Qv=ycRIXKJG_JJu<$wJd8(ULA36q%+s|0_Y#h|tI#uMg z8#?(YlzLGSg+Xm|z)q0o8nJCEyNavMZUI|oyJP8EP<>C9d>X3)=Ct`=l*1D%p`o=- zw{Le@VHstmPo0L5%BBRB>q15|8}jp4iFy-8?y*<*#6c_Kz^*%ld%98+Qt@%Eidh_Q zu>KUi8Q}~i(^lS{y;4)1CVetHln};_Osq6ty1F90VCddAp&|zdG|`NK)T{2;lNqbahtG-IwS^EmXZ0|%Pd1*VCOFmPV@jwe$hB@~m!Iu3l1({} zeL2j023$t39qsV6B`fB;*`i6cX5GzRWBOHbqFN2IJS0?ceL;Uj$&)Ol=hpLYANHu` zqi=sJZnn|2;ywy{QsO1K%v<*jefit6PLdh{CtPAsNI0KlI-L`~KT+{i)%2cWCS&K^ z_D<3nB(3QysT3_<6Y>)ox37s1x~OMgOiW|OkBOC*qI-(w_@0(MdP$*9HOA0Q;Bv+X z4zG`=Gu|}gQ&&VvQdaO$iywOw@R{w?as29Fdf8OnD$#7yrcEs=x&|#F;)zWR8Mu#T zPa0;N?ZzFc|HqZ9xQK{KXQx!lXj9UdI_0>*^HOmEGTyVGz7rtmlj$|Zy)eCN3!I0^7`oy{dJymZ*pdEJv-#0VZoe^o|0xBlAJN?;Xjt%J zFP*WJ3*pc`(TJkrPbe)yY8U^^`mZbkCS$6k8I*SkOi^w!axbL$C=sU~+3DI; z7j{%X@K8yMNiqe&_tt;bVUSzFt|GF&f)0aUJ<>vTfT0FdteYP)d=N`1j|fJx=Ek+q^#CGTlH2ot@XU87Zg@^zo6Kzo7$ z$OldL99U9$s(*$ue$_NVPbUR_^v87gMaTMFN5+bGB8?FdN|i82HTqyq(^Cl#y^pNW zZ|CzHI(I%cRU_%(ytnr+FvthT;`iHzEaYaQ>>i)k&Y<6|lPd87+tWu5Fse1CI8{sl2V7CJgtHWT!URh!G8TtH2Z9=V>V z{5OtCIPYts4lbDnij-ZR)*IeipCA*Lot&y-5OPaqQRXNul_j1w+AyQc*STe)GsAG| z@QzIh`Ig>gSo%iK=gVWeidzS5`-CdD=klMucp8CTcbk5A;NS#ShB$atdQS!~U;(b& zz3mswdwI1;qU6al)bjWA=ByU6m}3jB9rnkK7Di}tJ=mjSVRbsT?Vh_NOT_gK!?MrV z(|!!swnd7Dn>Q+mI4sXO${sC~oYpNYC^7n!d?vYggfY-tRU(=Fs(WHmX8u%+(`PdJ zQQ4adcJ7fY11RM>)KvkhY;&-c$`UvMP>9XG$1RiS07*I$c)uE3($A%aY+go``1REy z)G+wesvPxcG)Bxd_obfSGBTgI+*@H%L3Y4Ltj}q;h2lYkwPQ{bpsjsecUo2R@%SBt z2iR>CFVrB7zP0hNGWPN)9?DnSqKClgH}~r8y%0yAIhU)L zv+5~iZEN5A(y2GZto4K`R#1^p;#I#IBCciU5SNA<~mF+ZTh@803fD|Nt2 zWoKoJsfU+qjwk2zhSb>D*y(Kz%1*w@h@7J;lxO{DuLE!wx~ZYt_9WiaDo?a>O++Bn zHCO}?j86r zk(Gj0X-TF?VV-NFn0HSEE)U6#j_VbcmbRBjup#PuXJv;kA1rla`<3?eOIG3apJA|$ zh5PsKJKdb@kV8>g z2=2n8@-GYzf0aXoyGfC2OX83Z>9R67SN%AKnH?>8xG*fgL%*0au4SQ8CGXuqEX>mF z*q!F*l8y7XlPE3Up<<7y811fz=}+BZ;(mgTlP9-x+N=l~u=76l&8UyHtsTs93Lp*i zHNk;-L8s|`uLTp(d9}vI$dAo!ZH_Bs9QwX~m)ajB%3W=9{+~o*%hy1$&oCbJ)JUn; zzSDH;eUk?JugN02uDr^cQseU zPhfGI-Q92JHq!fbUpO=x3&SNnVbcl<3XTwOw*Q>k{^s~D|?%EyY@vh%!d?7csxdP zScr0e`0!!#_wU}&G_4yT`MCBjyyg$aQya$D_wAd~`}glH23I4ev9ht51Pcj5hCWsF zE9s2Wa}~jRzsk!u_Rs%kG~$*Rn|1xgv#`vH%}H|>nO)C%Kl@S%nCsaGU=-!f<~ZHG z;n!%MfWpalBg)H;m1*25W>K^G;hA}*>#htvlNq>bQzifJGyQMx2AY8czK~hTq;&(w zD|@u`me0Dbg}Ya!U(*N}zF%4Z7GirGQUpq1^W(olA1IHLhJ9JEcL1jY=vOO>*0$g% zp_J%&{I)P_xAuH~ld1D6%Wdx?1e^Q$_rjWOkUD@#j%ewRYUAC$eK+l#nXN5~!7uk0 zZw$!g4|_MeJ0=vh*`5G=X)E6c@fU>wi@m+1n8--M!s6oo*oWLV^;o<;XZ-JL1PHJH zCXTZjKJ!n+-yRHG=GrcD4A5R>G1#)Q=lE=|8>*9g&yr35>>By0!m+KhvoqRb-DDG2 zNkzqtvGMWdnwr2yt}r&bt@Jy`b&4e;k~6*t4vfb=gOt|+M;2_Y^N6fx&Trmjq0&55 zYtUR*Cjo#+7c+D7=79mCQCFU4?~}}UVYZ~?Z#unL;(BmOP|&GD6!U7;?w^z0ZS%#a zGq$q{2lw0eb+f^YDX!Bg~d8Sy>jVcDNU$GN%#!f@iLG6Cf_euVan@ z?6kQB|HmYf@ZTc3}csC=AkzI7ep=npCD>%`erdvQDuV&F=g{ytD9MgDJiFXf3{2b&e^+> zi@LUOvnS0~4<7u0Ll@3iw<+ENJqrz-px>*dx%u%4-?7I&HK)$Oz;%&N&M&QmKvuQV z{p~FTr?Zo3*DFqM-OuR7?THd>zf zN4VyFQSFvj=`=Rx8E0;@!ni;EFkm=Tsu^6d0PEkB+%(prN^n2>4RU9Z9|#zoKTkd7 zV_lqHb>o)H{~5zZqf(VmG4iL!r~-bXRWYQCaxE<_nnz7-HP`)q4c%Y2E2$f?X}t*N zZg)rC&020eerzxMgVx)es&#_w=<`B}i@lQL&AzeB$vYkftjJ^FeLZ0IyPdL90pVMz zqHJ#389*@blk0{5IJ!ZbX?YkGJZZy*4P)OwN{*#PAbWSUu|_Nzavat>(8+SCgGVHQ z<#V!5^j%WTzAIkF))uZHsk%^B!ur9cT!b3NXgTad@ZrfLh0ZavKa~#nL4-k&HPE<1e86 z{OaZ8<)eBe_!q2)3KFQs?Mp&p#+P9=m4b^zyN!3?4|I?no(EBiKbR<~pAhD1??Sm1 zZ6@VzaWSu;4T3DSveK@#TYBSf84X1wD0Q239+pukCqIfEU5W=l9XFM7zsU1tC`Q$f zcT-rE-l!NyDGhIY3KYY+)>)-qdUV5Rt0c8*(3s18IBpw%vSPu5z zD`T*58&;AT*2aP>M)^gq4g^vS4SD(--$ijxf$JK57Sa_Gl#Vzfvs0A7oj68-=$+Y2 zDG@{CK7~=Gj9-6(#O*zDZrZ+HjIzP5b+G>diWYtlOeF`m8k4Ic#;*A=%f=ZC2t>Ze ztH{tNu?xQiTX_6kRNLgs4@-%_N1udIDOxIllHG&xHl^MYSY9Z3*)L?NCSm^W3SP z`^RO>(#w>Kzg;YM{R00iY*c7l1pr%EOiR4wH{AQ&zz(`^OzYQyhob}A0PODuowTj+ zPq}^R`CBII)P3RK!7b+}h`iVkR4p6&%ww;-?f(km1t)xq3k%={dgaqy%t>>mT3qiS z@g8*?#YkG(vb4vPu)SyI&hYR0H|_VFC*>z;_EmnWX(1hV?uUjFvge}GxL2vLnk`Rr+(MT^UV8sE%$xD z?w8@f?y1+07J3@2-a?cI@}~ckT3|yCT|1_!Sw$bj`V3m{pyY<3C3EUO4pVZK2e57f zCmAY>c2=ZHtf-7x@O#BW0sk$tm3a#mXownm1njG_m@JAnWQi4|cfa-YqG(;|AET5} zutrY$q2NH%a@}QTtIrAgi6(OV$GExf*B&FN=By;Tm99#|xd+JUsfrIjaGC*BJ}$2z z@kvZn6t5G9lvv`3%Jg|PIy&mEcGI@zAp1;;GPN}~2VF!l^W<%TFJF3X_6sZR`AVq! z;!Co)yzgrsZy29<^D(7eZmQ32@@3m~y@$3nP3N=Yuv>cyKCTy4zO>Qz824qH^Cteq zANr+>50Autmwq!}yLmWZ(CGH5iCpU1CgGlW*$ddfSJ$!T&kxcrKsfFENu>LY3^8y` z&u%B@l#DcK#+i2!G8ONzi|+k7`z3sMOD2ox$tK?eq%b2=zY^1N53wGj?#&}f8a0Un z%@FeJv3%dXD-zbyhLNw8iguL@Hy|`Bb5S{bS@pim%;2i3@X-eB(=SJ$r1Qc1=gQr@ zI0?trJBL7p8>&F8ucvJeqYlVZ+gdD%eir7Pj}=`;((-O}pUp9X6`4v=B$RG`G%j?B z7bCs;UV=mV&M)(YLn0&Pf#$2lD8uTJ5jpzi)oVjrG7;nliWqKGP}3$f(K)n*SaDHr zK&Ywr>{*h9S4}~2u>Oec60h4e_g>_GT!y9ETV+62ta~*_$^-mEeSMuAyGpNBIL!<+ zWMOX=JxT2ro|+H}oUxLrSTfEK3MuxIs6LA}J>cWB)J{}Pncfl)q-Db?T1K&;$_7%K z=Ip)DtIcn&vps7@UH_?z_l*RIFQm#2yLT`A!tThDIMv0oejw$bM_sDGnQqqT=BJQ@L_sWoiG?Q!Va4VD*nZL=wV8Pu4 z1bM|qGn$EA=?Y?sv&@}Y)7(2Z*Aj{0(M%Qdq%XCQ=DLaxrtFAxPBsZ#84JW#?S25h zT)hQ4#TkdRLhnEYTu&Yac?Pui*fF*z)PWQ%5KIDnYSn0$9-tu)iZG1I!*^M)5dNOR z!SDBCso*jps$5r>aEs;3k{oI+5=g$B%z+-qN!}3)*>4%p0wY&_yy*e@wQ%pr zB1>y)>;*Xaqx|A3tJ~abInIXEzDxWxHa})|?0vv`O@30M^=T6%QX&FW_zE|1Fk1|& z!o=Ts2dS&8TY|Zo0$qvypmy`lKC^62m)^HyL3p?k5ZN9SUaJp;*_U6Tt?qeX#rT!0 z^iQ{HOZbkC-<}=l&rHA+Gh8~&FnHU@r-g8-uhzqQcd^f=b!ErBrJ#v}^E7onRt2E! zvR6|SSLhAcr8ZrZ?wh*-Z{DK(Tl~!z+aD#ayfV0A`UD7{`H6C!l-&l&va4zp}+Jxv{5bV!4Tqj!yP+ zs%I1Mz?iDaC`yw!biMYGpI^5*x2NZrb0+QyDGo6-|NXL6bAGx9!q%4ui*VN#o`Kuv ze{byzt$o(b1ha=v=rdX(%4qb>?UTkQ65NMIn3Rb>pfE)4Sz-^OlYfY}x4PxktDk{7 zKO{Sw2v8uVfS{fDlE%3`{n}M&;kO|Yi_5WKTE0j$fK?DyM}S2~GUS6-QEyZ1T({dG zt!!qh$w^}c7^JqJ^gyuGMJ2(D&!qVCCCE&Y)`tu{ruO+E*{d6U?}vsLJRw=P*&C70 z%{5;c`fb5-IeGcWqETi06ccuvF zwm*_e87^SA&S1tnf{Kf`g}#6VQmf&HZ(4$s;-_cqzjeLM;fq*ofw7yBxro#-P#@bV zF9Pt=t$QQegpMSKBxI+aI9mFndNjBxQa~}Ed<3e+PVg)JvXuJIw2AXA@$nln-y90= zefM>CVFeoN+TEUW`}kyiq>d~2rVvoq3Bi&dM`i1*LZ6*hs&4)x1uDnSR|)YY#=du4 z#eK#Kdb5G*Fm`g&^lSY`ke;3-?K#@VXQ&pjVaN1r)?et0@*+2PcR4ejVU7_`I5Q)o zO@2SWlDL%c^N>~yKW2T7Dys~X3ci@o+xK}}Pk$`lyxp~TIB4MIzS-s3<}*SIK+v+Z zv@CtD=vaEq-hKzt=DVUgU~tJ&$Yc5>z}eFp&QEG>5P*b~E9rJUfD!f8NJc4C9T{%U zo(Nh?GZrs`D#>Tf8Vv@@WFpdx$DtVVJ@ddX^>aba5BNg1b21YsQ& zGxT^AzCwPUi97;0ZJg}neva}FsR)LhG(~P6Aph|=qvz=46CDxg7oW>o{FrHCdCvMB zGXSGN5DLZV`C49<5k9EX;52EDLq#`aCY42wLU8K}B>GHEZxG8Si#-22{N zWhe^(UE=9FA_e?neCj1X1>bS4R`Fjo6LR?v8p?A~JO-3*Uuqxty`c__LWx4wz56N*1STuO0 zI#QJ=`}AcT`QaNd1bS(OxQe+Ns}E_{54^-ruh$p~@QToX2B8F(I{rRAP;+M!iTVPN zE^2?I@S``Ozm7H=D<}iBWMA4+;ew^*Vyzkvh$UY?5m}JbSDwtB!0wQHv3d51960t? zgzm$d0Q!(M4*1X3A)4Z$pT-8THO>JkIgd8!!hw_b$S3vFznqt@U%E=`PX&Q8I=ZAz zX}dWN2uK|<+z~H~?T}aQ>)VCG@KMxVr99`w(mx-y6PFNt~K&tr%uYlz<)al1&5F)i9BDU97~xVj8LF`fxD zaEvBbrDas(CR%O5ESY5!mP!pdmmQ+D_|znWf? zdpK8vR`Fef7RO$_ZW`Ku@=J9*qCUSlDmJ$F$N_SrU3*?;RZgRY&44Z#x3aM(?Z(I3 z_nt<XjFBcEOhtSB_^mA{B%HE*#lRA zDnAWu2_-d0M!NzmZEY*|>}h)hfEtE{GI&R#LS$#h@x#*XUf~AMUy!efU$sE>Va+xcnC@@T6`K+b7?jboQ)YsBx50f+xj#4Hq;(-n=}(B_?K z8^pv85>heJDN&(s$XZZcVlFacSHwR?X+BHVpw=+DgYmYJpk-Dj`(>;_B6GzBf zfvVieoSP#zM`mKfxx$Y|)53b0CMgO^udSfejiAfZekVyCngzF0Xs z7KrbgpDRpuDkeWckPRcAKi@rM77Gg+Jj&7k;6Zed>VZv-2+AqDx3m>*(BbTI03&0@ zm(o>Mx%{Df@2$^UDW8DdA8^kA8@-OLXJzwC5p* zE`ykNEyXzR*>IB%Z(izJ1~9&s>H#L&_VKW=_Z_z7VHwSHrvj@4o0lH|ZD~(ss9Y)P zk<7i8=4N$EGH#V4cS>__vf@#z)ZXeDbyK7D6j z^l4p^d{wnnss|4T>>-B77dg$YD_<1-OxucmcJkj!^f%toC^6?7V@tsZqI5+n87FqY zu76YDT#BgR;b8;<5!6FFE`2BqBDm*4AbI)S^fhPGy$=LO|K!t8`{vS%B%RImquM$;h#j4bu$M1UmMGLV;`#$vK~`o< zt65RhJbtVEGwmv9hq~3{s_UZd`}@<0=MJ@MeW{MHxkrEba*N!dL(fTZLZP0Xp3W=o z%^|Gjd77N~Hdp-W{fP;3;Hd^`(qpPcMnz@2X;WBTt=WIYeJW*13}beJ2I2l0*v?o3 z#JC?{4|y~~4K>=f0*tbX4oy##{kJUS`K%NLrHabR6#TVtxMjCB)GF&+*GTSEy4lG9xpG={YNmn$r1Rp>zR9<&dv=7X$pb>bNt$`f*_k7 z1QRB$&P({M--};p{rX(l`Hc+?k1iNLulVT!Vd&?3@zX!wn-4c^1}<|iyE!}&W9=<| zym5T970!bgbAA0F;<3eiT;Kl`q@y(bLx5k&L$H4> z*C@g3{hS}_u<8D z{W{(#o%TEueSj>RS?D1ZFD})C0X~TPYgA}Siomcy5yiTBRaFKz`O3A+SK(>v%Q>}6 z{=K(Cn91CphIl>YVplFRRnrTU@~NTaw9XHNXUH*wBKo zZwpyV(B0ls@cseI{L%VZ)z*;i0*!E!Av{>=IZY?C5+PKN-Wue9qfa={EJ*W^Z24Oz<8Rc* zc02Ar{QImMJbd`bi&Kw~3|uOX6~MO1RAvfoa*69I81uOgwsI*O&84&}sJMg|`q!rA zBwK%a#+Bcj1MHx%Di)ugZypr1?xY>)f~AsL;w5nVhmN}Ds+S9LDTrM4{(2o)UT=0+ zUS2p;eStyyM2R3%+F;jv^o<8+W6H;+PYwDs7Avb^7u(@rFR(f&bOzXy+n#qGE1hIc9H5&$ftA#ZaHA4)ys8n{O~By#!`6@2#g^>a;aBR`PR$OT0PC6tLnzSaUm0hV-G|+NR6I7W?YSa5rds_zV$b400$<1?^D>Q0`Db zL48)*-iZS{7S__iqE(oEe&!Ng*ud)9?)uIfJfA&hT?s#D3LFG_g7H%ZG+?iJ^U6sz-?^qU(%y3|0v zsE`Os_E6vi0cn?`vZ*nOabi%*%C{Bx`)~%hFgs-p*Tci_zXBmj1#D#y*mNkdJTp%V z*v2SF+MX+3KD(&~qLqi^HEQff4iXiCZJDj?uaGzf(t;=!3IyuGG%Atp;&XUnC9^gB zNda9jJQ{fG-rrcn~`THBVVF!U>+14K73L{xwhnEhJu~qc+RJjc32;Sho6VK26e5F}Wg9<(f zw*;tk28k73?sUO{_YIf+s&sMJPObj>6P#y!KWtEOLICSq(`}c?ZV{oOM4gSiL+Z6hVkA0&_|TW9Fsr?HEz31tqbd0>Q@8kW@t_-w3ND=8F+^@PEr zQ&l5W2?69>RgE0Ztu6I#iUhueUDL1$a&`KqjT_|w&KsJY-Duo~_yb`{RMOB`xnxOm zQGd3Fs-Vs+tl%qTzI8={9`R9r>E0Za7f$`rjEpHj?;m_{R=F@NjUDXBq-i;+>!PK5 z?9{S>R{X0%r*j#1Z0u}~&n4|MtLIY9mr_?-AK1t&yAH&NnOTRvOU18m9vED-=dOh7 zPMeZ*fV|wg|8F|@tO=EmpJxL@Fz1NVN-!lQ%N3grJ@aN4GZQ`eeq8&EjE1Xvu?>p9 z&$>NuR)wA)j_C^jcV8~H7ypJ`x;x^28BPFHrnU9QWkA`YFrW^3bfvF*IQQF0N5A#b z%vQ*wYc3T$ZO8L7{otP7zfRyZC3SNx-{?*2uT9#}86aM7LcVc;I%}2`Lj8IN^}{CUm3<5NhI>!f$X$7SdAPA>{(wb)IwTkIVm zojp(o>Fl+$cIq$kjaw`3gSnXpz#H;1DEX=e1|D88Dv0oAH4f*mauS}q8R>=mqM?`< z|0e$a`D2tc=`pv*W}6L8iO%n=o}JLaaTr?g4?M9NKcX)EXy~ z+QjgA$NsulJH&fwId}I9wkO_Y%+L$%s<>x4fuWoH|k z(mm4E-LCtkIFb-Phc#)Lcin$f97ds2J#;b!T#cl`TFm_ zv&Fr10>-kmtgLEks*pL6%R3P7m#~xr@pz{rlmZplMO03pqWi?PMw1t_cBy-s9RC7H z&|r5r&J+#q@AsZ0V^pjI<9vWR+0XlvH>!+&5bn{0u?ns&ti3|(2}k5?r%0%>9^`p2 zGru}%9tye%Yz=>Z_y9H5@?tH-@_V10?bg5_DyeIQic3Bo{GAYmRJRj!dN`PYtFTXj z=@L48m=l2)-91L$v9(@tze$*nWrXM=v($eR$4oiUx+Hi-%zho-=R!86K+NXe^uDoa z&N<)r*4pu!=gdG?dL#C$x}&H@7-=bfnkS5L?GB{&UP*e30SsFNUb*BW&=7O=KX5qg zvP_$t41q=1>}Vk0&G!-p)_Q66|Jz#cZvtX*6_@5mkA@&-;^E-V)ie{cqbS)@l11~M zHnAkR<8I%L1{1c7v^myLHgRm0H-~*{H{DzDAIt7}m{@$zG^V!q| zzrbJc9TNWXX#yoXhw#WGEFu`rO>v!Qv2tehIq%IQu_<>sjMzxl)7=Mk43Tfn>>ltJ zs-Qdwd;YwYDVh8DhR;sqs{{fzKur6{_fzkAQP_IN0Fm7Zc7FI%3xsm6C~=ys!5;Ii z6ork;X&<-g0i#~ph5TEMK^(p;P7t)@-_7TaH1w+H{Tl&tC@qUQZUDw{>gvp|imjG~ z14sQLOOG#a)m^1)_i_B#%4aq|sLkK8;ai%U3mMBhn%S+Pp#hBU+T2|(Xl4awZySj${Z>{NE=4iY<<&*&WY9HR z8HH00kl_7fr^+u30?1>zb#;GEq$fnbR|3p+`5v8KxM;ryZV%^l3uUX+BcgxB{a3iP z4K-Mvi)2mQ@bjvJN~)@R+vDLmAJgsB7;PBoHagY$xJi*)_DHrhQNX`R!vKig7x>)s1yHTx zuNNAgVEgD?Yn3ZYR}mu4G|0Gdnki?V#*n=QDe3n<(+4w|UhY0Cf^8(5j%DcWoZe5v z&ogO~OWv@{4R@E>6aVVj%lMhCGy3Owxxbu0X34_2_G{=qV?nIPE9p^N?#$Y*cvx^g zQicgk`nnOG^2+&t#NuRUc~X$<73KL=*^xQ&?+2{9%4L2rz_Vh4bU+~6g3xNDUr8gW z);Q|I)$R`L85S@c{`@_==CtIv0e4S%uw1=|MXxD2e^8hKEUlWQ*{5HVuTG?=XK?TK zBaKhU{ie7F_gp+ZSBp_EqC{vQ|`b_&Sb}~ zj9d5^7)cMkBOnw$J(ONw&ocb0#kVuT+-*uPoO*AwZr3&1Z!_ zsMlNT+0lV&(L_mRRQ>5=M}A2?%uh8teCJ{{0cg#2lDizw0P{42d$me1Gd1OV+dupp zr#W4?wEY25j19bsTl5AG9??2A>u$hsZV@QF|7k)TwEd8BQR=5N+lj=_pJg4tynT79 zO`;mE@lEV)&cLnyqvu=q0O3-7Cl-YIA>7>DN0{tr!IK z6<>+s2Yflbq9EVhUvg3Ya%`NuxC}(-pu^`Fu_u(7L37uU3})B zxSEu34OD?|Laf${*Graz@-8$MBB%dToeFR~x~BS`{RWd|XH+bj2mGoa)aRsTvP%sAe6w^-Mo+f{8&e22! zFS}zuLw6{F^@ecCr(32zF4UE32M);Wo3kB|hk>1QB{0AM6XAftN7ip@@GEkdnjB3I zsXs(ES(MQ*@uM61#2(LM7E%?x*gxpK-`T5z$A`N#C)d{IzeE#n^oszciuMpbgH6VLz-If z)eV!Iw0FeFHl8;!kE}R&vd48_O|pt{23>SG+teY@wmu*rz)1gk1ia$4A}w#+#=$cu2n1N+U2WQg+`f`x9lx|Rd82mtdl{e;5Cm6ev? zqUWh>&mMH-yRbiM64yTV&*JWqg}%=Op{Yonyf;Yw2E;$8%bE5o=g#KLTcGi8cmcq8 zfFJnL;V@IcM$CVmH9Z`2!kEv|hw;s>6E}W?wShyl&Zr_JBT%`$?VEA;)Y_)JJkgfR z@vpak(*jUDiFS$M4}W<35C7f7d%rb5ja~W$#G-F%WaRjcKAr9Yn#el)ho3IZRkR;| z@!$v0!0Rr53V-+9Vn5YF31L@g3tYnU$q;mkQedh>_C7~ZOquiE(M1n*`1aj-DndKv zW}$hW`ViuK!KNmvMakk_WwohYcE0;3pqW3q0dUuoS2duqiv{cc0l9SkHYo$W?cRK$ zzU!i3!_(Rkm2d?90ElPry4PZdj` zvBsK)E2m`wKO}mV*uN(FONy!-Yn+9pD^vfhs(vM}M*)8&(zWM;Bc<(vBWt|>wcFX> zi@=m|m7Wf&9Aeuq8cyZ)M#_(ePbGKQNx1|<%QTW$% zJLc@xhm4`1H^2i~^S?yr=wuBPeY4S?THtpq=9G14*+c$J{xs^;a2J39l9I8NMhH&f zSzYpS{8|SLk2{(amOM27ot527Ad`(aejiLlC?~tbW@`@Bm!v>SsDJgMni$YPv74^w z>^b!&oD2Ul-<8eV79Fe05)%OUjnYcv!?Pt%33A5j+X$ zb3@)hWlD`arzpSN7tZBNj?t1y#gQ0*4gk~1L+MkwcusCxsW^>BV?>X+a_HF5QHO%W zkC#(L5d=w-qj~bu#&0e`025)fKNcGj47}JQ58YHwX@lk8Jd#s5NOWdfW#F*~)pGYm zKejlduc=E08#ys8{W>4tMlR*-5W^Upl3)D21hhH)1-K;;uWhS-4&gZKkr;gqw`0NR z*kzuUz38aV`}|e*2X;5Yu-+>Z_v(EUQ#}4RF-0OdUgm2_{njmUm8aO-V@;Il-<^kB zFL;Mge-9pRK)%9dJKoMR2{C+0Ld4!?KGd;3k zkA7UP&o+0%`Q+b6q6)4wZ(I3mX!%Ih- z(3U-yO|iCQ%#usP99b`n*Og7`2AX{WuFoLbmTz*d+7%0&R?|r?uR8$`En~vkZiIP_ zs;TJBWg`d3sNYNr#2Z?T?2~G`Gj#xQ``5g_!iCv`N$k@#v|~-TNE#`8B_>~xsdfd7 z-hOQeCPyp^`Vpr%^@VG@jkD41R0HzQY}?Ac$I(E}^Hm~)E+ko3uwQ>l4a4mco8R8r zDzFgD#4sst@JeU;-33|r#G>OP6NA?~-98f+)kO?w#5+#;Ga`Y>F7I(F`Mzi{Cl&Q( zWmAX$yCym?;jVh*>l==l^y1-YvX{!r$Cig0?SOLzB!c&64z~&j;itOsJ+2qkKa(S6 z5)>4d?D}3d+Z=zp_p)ts_C+vbz$~1NZak>|F5%C5&tc>q{Z;_Tc%mt%0&u`4;;uzl6gV&TKqDO2UI}$k|aMk{DQQ5Rdw7e zf#tZJ=<&YK?sbAp5qNNnOp z1XoFLK1+ncc`Sl*!ANl%OM`BtOBRcV&TAQswu*F^L= zSAb3&c)b3aX90L};h#V2I&xUOol0TeY*}6W9*&<~8+sl=N4*PUOo;L92VkG2GT2sy z7gie@iO?M{y@i@*4jUO{Q5Siz4}TCp{n~%;AZllBXdI#EO0CpZLcjU>nsdm`BL%xYW?K-xc_t)xs{1U!fm{P#-*+wO);1^=YX6gs_fWx2ZYZG)p$j35%L$i?~WYWWW%$Ovv%J|K7GppV_;R%z4`g<7n zTyrf*Fy2sR&2}GFjyS9gZ&%YKf^#^txVpO`w-+Sm8*-9O>?_YVQyFY9jrEk*?_ZuR z`a88LrLP#z#TVq8q@eLA#=H_vI85y*=k?ViTfPTM*#$y>@(1!aCx!4viY38fmn)7m zOBH(TB5}U*_VmwqNzYhb5gxLNTB>w(_EZ8lL8YU%*wu!^pW+e3$T<@bBKsl`${-9&uH-FvD1qJ6CFt<#iul{v`;AO}`~45I z5cBz+|ANAYbXOo{{~sGZaQY1f&4k8&u%m5wy-uV@<+W=`(9Q}A8mxZ`fS)E<)O@ne zhG&az&m9}IfCfT=(Up+ZE|^SMNy)t&U-9on<;KI`$Xhb`w~ZG*O+nKj22Oj;BB>Nt zKgheR(P^~}yU|`3Ggqe3E%&E&o~Q;&`9N@8>GFD=4QpbO1a^s&X#c%E1kMb=bJW*jZjS;07igP_IsMavDB8-7F_ZiE;SlwuLXx&OXVZ`44WP;I(L~-?XnGFbtZ( zdb+FN<=ld&-mHkVW2m=$)tzAM-01{)qdPsdrLUTB;XN;$hc3ETJam=zH|38z4{e3a zBF-Or{SfkxuZUWAa?Y)Sv45~L87}z&ATyG;eqClT2a|vxvR~I`T=?OaUfH*MTHkT2 zY3>AE9ra`9_A4tO6|Wso^6lp}(!WiOAFX)zZi|J?IGEpMYNpMGJvXNXFBxGzHb)@r z?7C+=`CYhRI1Lo;{du5J%H0fh{mTn5^;c0cRaNCXWQp2Vpi3UHF?Wsn@KEtP83&vT z9vm(PC+LdkE!LRqUKyECCQ~spZ6A45b}M8`HMk6P_s4Di?yCCZeDaS@7{BQ0r|{eI zWy``iJd8VnX3^plfy5?QA_Hukb0_yc9J1W`e#vIT+hRvX&TKjqR(fmUp~G8kuL-=j zy}a&@L;E4u?$9lV(|Y`#daqizBI!wz=a|5H`JJB2_6TPj)!nY)Euo#>GSxT0h%{0$ z2@cM+NhPNwI8CVb8r<5FNzJS=WnZYdXiIgfR%K95R(N@W;qKUq%nw&NEe*Nx5});t z>BTfEUYpQZli~kT?e%1m;&`5#jBIU!bG*qK?AkjG(%~$5`bk+)?$l5VjTGf?+VAY_ zoKG-ychYTc7xZ7dbK(BYbN-p+)|=$tbRG32SCDEj+)RoH-lhB8@$`6v9&GEL)=O8j zNA)sSh!(E8*WAFME6YYSeAq0eK3ddW>RO3O?s0~`uj_ZU1+D}Wk#);Ci50Q5R1BNN_>Lzy~N$#xc5hk_&G&@k7?|wUQ84l>?UFv5**T&eY`d6 z-(k^Z(0Y0jz*jz_7zO0v`=v*GuK*Hmh($ng>tP&=su5ockF#V+`Zn`VE_E4=iXOKW zUOh%Fqi-p(s(M*NclGwv&ptW)YBFj4YiH<)!5_A4gJFi!b@pOI^3a^%^Mw9zCZ49u zi($VPi5!BA>nQQDkqP|PX4Ru}^BNEG{% z0WM>E9AbolDaJ5(7}$g;c`}ni1e5;9RvSo+t=~B1r8=Ho?k9ulKS@IHCh$1bvf3~3 zh|Ju_3ng|k*|`MmgKqAxbgY3BBr(f=Ap-1~YNB*tN6$9obfF%`EOR~lo<-Pat0 z4msxeJ&^2kUwr>Q^Y8~7J`~PN-2!d3m*AKI)SGsRsdtz@G>3p0R`;cH$&oId zqe7(heWU6Dg_r?MqaU~QJ()uyW8%!n7HoWn)>PESF|U?4ZkaE${G48Hn()0{zp}wi z4-*m!8ZS}%r zo8VTB(F?N2hciF`!LP~mLUTWEGYF=ZKAJjA(XE;?yB!tA8ib3QV*X5kT~cL!6|pDg z_X{R8rtS^Kp-Hgijqn-*^;|{qted@_uxxxejmL!X%5M143GL(VW$qN_q-Gd3aWUPO zN1m#}xc1c#3|2AFLoth41ND=M&dU^5`LTAqW%=IUX_lzXsZ(l7ZJ|v>fn9&0X16d$`rrry1osrY403ygW(-(*BH|+HHRD+q!jjC?2xdiC=>~I32<+HcDsv zbY_}AZoQ&MjzRCT6yg+Ouv}kJ1QJn1?|B(=Tn@9|!%z>y*%eQZ^|^na%QARufR8EU zW_QzN2fe~OIEnJ~#O{v6-%sjsvrF{>y~R>?lp;nXk(vm{#97e$S?_8aIbAV?&TBUB zg>+gO1Vg)l)Yfd4DYuc+UHcf0O6AdNcyC-eueG>G7C#;|(!uBo4&vRP&MQ z@+jzPi%~m%IC=4w>_HYC*_oiaparkPOU_N%i%uFy_KmCY=#lf0yv21C(H{F8!)09# z$ALNe4|Qhn0E#FRfY{`lU_5JVY+g?PV2eCm=jd+J%VW#hE+$>K8cV!6^W)m>ziz09 z=Xtg@78#RY<8Stl4NKApJa_A6AzzMmYV`#S3x_WS7&y*j07dfaz#AZTEaS;HA5^i< z;C_m1Brh<$f+Cf2Gr?__rM8B5R9Wj`(z@NrY*PlDTfrk10gy_xY7n!BVFh!)I;!1a zTP9EDhnxp}7nBoMX@WS67L#Xh1y4X?n6d2I*}Z3{o|!*RiRkoo)EC%=UjvUa;pf{x zo8_+bf!KT`x#%-zEZ7T?pwp?vLAN#<^d_ktIk@TC3ZtXsI5@_sW_J!=pO?$UHJF_a zXElyp|MJwcsDwBC>E5?M{dO00+ZgMd)YC<*)-dfI$!PVg@z>m5HAbFL5w154JWjpQ z@b*T$O@$e?P-3DeT9V1OwF}!ert~7fSuZPIC$fQ8u7keWr>^k<9nMM$HM^sbBHjA& zxc8cSRTj{*-^Qx3{*8LYV5ow~oefsMj4$CdOW-}l$ z!76Da$lJH7r}z=Gyi{!9dF6iEsousB7WISIn=p@%lhopktZ1hs3ic#S+@tQaYEpbn z7?Z6Zp^U^8@VH+208q6O;Wgq}R2K}Cr!?0b72T68kG5$7zYUcCyMsG3dun@Gg&z(l zwnl>*OMltf0;uqkDH(c+#lkQKjAjdeB$(Hn8uLbt<;Le>z#*9&O$B?hkO-Z0R&j24 zyGB0O{Wb1s)AUfF|5+>r30)o4~IAAc{useRVlmU6W`P+EtUqGa`F8u-08Jf4_&h!kcbT> z6B`w1uu*T#0s~XmbaxJYED^TNmS2*;3eXpIXpnwL%}Z|)(fz1)f9&Kx)Fu@%=<=%; z9=Kcr@9o!4$sI_}lk$KbUJ+v=Pt_X_1~TylRGdIQ;APB4v(OwTb`Q4~#uFvw;^+iP zvmnA>70dL8^`V|+;_gZAzy^AqP^Cw%Xf^WS4wL-xQHq5`L9HsIO~}uw$1tJ*W*Mp2 zbwO@GsHRspCgx6kZ~f?%n}4uWJ~iCD4n6t)0G3)j=Gt#C?qups1q6;<#2Sz)!v%@_ z@O_Vg2Lf0c>?EWz;$0#I!_*Zqd$NVMl}gf4XJK1+uo^C7w4q?qGZH`izurP0=rrRB z+~|5f$T~dEWo=EkM&YcQZ2GHr%BE#z5>+R^6ju;X6Xa!Ty@V;SzvQ;NxTUm6)$XHx3UEzHt{JC(^NECYU{%o z#I^|^c-Mcg{I6G>|KZdFPXPYyF{LrTj1{#@(!1Uyj`KSqK@qp6YMT$b1v_zYeuSCs zgn6;Tr31@rHhVoy*Yyf^2I_ph4l9LmhS<8 z-)uTwxnyN!#k6<>S`*``Tf<^dr+I?`CcGmzit*#qhrrLzt~#vt+nk3x@jMpGi)DR5BRUR?JEmyV(JR8g9eixLJEkAJ^Y z=pkIXxNHw=l0>0|Y6$aAA$ubv2HP>bp3ZkK6hNzLu+GYhq2rX!zko%@W(CNV;Vc|l zGA1Bw{cDSzxZJiw{>;?IwwbsQ_z^tyOl1r=rQ{f=8dRi?fd&pC-_q&y5B>Jp`etH5 zBk7HU#piP_?xT2wOc*Pc_z~V+4c*4^J&gy-`_@0Kv5{^W^4a+T`hOv(AJJT@t_Dhh z`z3tkx78LKjz63GcW*5_b$5O^Yjp8R>J8CChmQJVD@q6l&^9xn+&2aKqYs6kN0wZC zNG2Z0j|S8qrjF)8b-%-JKF_{&A~P?Xr{KOz{Q9eOgLbHXZ`yQ70a^Lql0y=>v%m$H)U{dIEwWX3wEXZq4 zA{$?8Kvdr{3;qEWIb(Ur%Avi$Vsu`u*vUTnAi@b4!CgcSw~BVZEIT?r)6HKU|F3Rd zGU-bM#Zy3h2vpEqN^)-dkbRp=fjd|nzD}D{kAHIg{nF7JKfEVe%f^L%<0oFpOv5C-uf7j6I@<}#$@Z(jk_QlKeO zR5v^{bThSRV!%9VwA$N*xb%lmRO#s%L8*(b@Ig$OQqWu))(Q04Qn4A<{6bC-CavCT z808{bc&Y&EtHz*MU^h~&XV5b-Ts_SHvb7$I#!?ET_iSVt1Osh*!34|OxZYW$UOw}meY6se5)NCz5yZiI- zj~#?UNlmW&+L-7sL+)qvgsZ*8T-#stV@z(|QP2=yH))2F@!SCTqzL+Yfd{ z?W!Nx1l98`-^veTZ18&P^93M**E~F86AYc{=n&O08eo2%gjC0`Gk{m)ra>j}@e+es z!#l=EMsVoqDP5Pow`S0*>nXqf&GnA=ouo;mO!e<&?w`2t@A;VU@@p`xSd2>gJw8~t zN_ci-d~JY%u-}Y1=d%VDz5car>|~@=mr9E1Cz%q0H390bHmoQQ>0;3X*z<{H;_QYo z@6bU4W`qCTwU^12ya+N=fh2+(xy%5}-Ax@;`ngG^f~8+g*PARx_n(CG@t+`1ug%;q zLR@vmAf9VW0;BK+jfBZA7&nz<1!K2#Pr{~(dCjV{N;#6laDsfyXV0( z+G2)k_gbup7bOudxl&_fp3hl&AFr zj%7-kAcJkIF4a-_)E650lNk(Mq0PZrkdP@$G3RzlHa6h9{4z;h{=84-j-0Ni-V#iY zSwN0g7|W$@+hPySAe?IIK1}B$P3PegZ@hAWe`P(%0MC*nDd0w}NI)4g94&^S8}JI) z75FK5S$@lzq8Ez=#DlvQ2fx<(>xCD-{8qIBr{@*J{nE)tY1m_(`*YVREJnL+eK-xi zZg*>A~Z^3k5TU z^|3YPdSn9i!_Bk8cR+Q(2omj;;o}r?X?5yzpSn zNSAtQUcb<39W~}XA~scxq~@WT$IAtc9R?*TSJ)ndIl=oBm{@_4NKsbFgBAX6^`>2y z_G|^BIAAjI0I)HwuRsValV5>Vcc#Yr+1*At6PQao~jhfpnBNaZ9t>8G|3oz^q%@chQKvfqgP74!EgF1DGFP!%lZ z4n(@Eq1`b#`KkS~gYF?zG=t~yHTzMbdo{W;EqFtZ3nfoNy382Y27A)(#5o8AbnFVk z&@w*QXv5^C;LyE}SazG{gpR5u3h&X-PRV+X0)116{{aRP!0E*)=QQcDT{`)Z$))p)4#X??@I9F1a*Qh_(H9Rc3a4?m;VTDPeS8|<7 zG-&V>Z{9B^kK{*u97Hy`8s=UNAPrXc=NXx>ulv$P1rj|{YveZlE>P{2z070z^Hki4 zxZGskAS%r*PTbTrtDlrsT;*jJ>EuTk9!SMk>t?^5^>BABwqZyVb8S#H@J{w;5m%FrD6I2y)2!6gq&f z$pv6)4-rHq!D0LCZ=AyWHCnp+_ot2me9xMh?DdyUa_!ipkC=v}92)pdNIl zEuI?SI-+FqufpHtYbqzi6)J?Vq^ z9j=^H&M$P2{tOZj1Oi2U>0}`$`6Ms5Ky)&?S$0`lry|)MfhBfvO^R6Y6cQ3~6R5K~ z!Ut*tSVcY}JH(nA;HcdgQjfMkqTGHUp`+{ZA?~piflplJNnZ>(x%7i6twe<}HXav`XK6o}pV@yE!_|7_;1e1uE`AQq)ZE`YE)Rwo|YET|cdV za)176)HGj%BjLuiP}qi4{gG|{#uwByfJ^>hiS5xawjN?Y$d}S_P+<`6>vw?e{>qEc zcx@}^br>t!b*na~gSaq!sO{!T6T2RX38#rhQOQIdE`!`WqVw(24=?sJf8h-M#*O9x zqFw8gs33I0y`cW=>G*gFyB;brx1J7mB^OI@yE@;!LV~cNGr)e0K6-RKzMLBdHE2~G zdc=Qhn*!DkNkma3HXN!_Tnw8bchq5cM!P^hs5hKJp;z;sa3}$ScJSkiJo@|N0$k4G z;g_yPMh&%%-)$NYnw&_j-gS_EQe|8|D2!FCbOMBp$OX!k=>Q=s_bYD4!lq8Y-$!M! zEeVRvmo*I^E%@v3T9zCh=0-I%)Oc{J&2$D*siFE}qAlQ@u3TX|kZH-qIe{52Dh1Db zNs-`ow3J}PcojYT6M7TN?*$>s@xOaPaTb$FL+(Xir;k!fM6j^&qF}%QG_?71&rqs6 zU-4=>Ib%@1Nx_tSVr{W3x_~|luy++ru@L!|eN)ZX#_7tc!JLpri7XSiZG31ud>czW zPqRu798F?lO~EJf$Y~B;VykIU=sx2piGqH*@*rVTWD`~K4Zwyq(RW^6>|60yM-{SY zVfci_AUju^?GJ*Ni;IPewfUP;7Jx7*I`W(|_Wh0i;b^&mH1h6?NW-f_t`}FBA_fQt zXaYoOll)j>OUngt*;-Tl((UuA*`?cma~hgx?gwF>$EZaf-WbpZz6pa#@(k<@Pp34l zq%Q0+F9|1yI<$T{cDe5u7rzZ0%Y98<>_PFSO)#6tI=wT4RJmaxG;e~;Eg-_q6!dFx z%qjtp52R+g1Xn}UfLdJxmGH*qBxab4;CR{jhs8EDz`Ie6FYnZ>a2Fb46-=`43_vxI z1e4^+4J*zkpfWz458NrX{TnyBPYX%|;MC|e4VY%|s9c6$6Q0=&?*%zOH_rt1rsEK` zfmcLk7YUgl8(k9}H)ZE4@AVSW=5!sPxf5BITuHT2YsM!j8Rt~IDVlC|S#|oX{3*JJ za}a+m*Y#~aHL*-Ja z1*xcEok8g@JHHi6{AR<@;>-JcOMM7V?L~*A3O)ANOPjD<80r2dhAMZ*?pQ=9>ibSN z#3$x9#IU7qHhK?*HtMaxmIONBNDlUT65C(s-+A2X#9@tM3~kO&0GlW&QdKyWU%RPn zJnC0^nDJ)`i z;PiD;g?)$pF+jaHBu%sk`XifHvP0wQ=`9bNe`14fjr@1;JlLb5a2C(YlhVA z!hy`pUdOLq>&M3^#KmT}K5I6s`uwJAv%s=*tXJfBf>&fb2EW!)?TTaEHKir{E-G(2 z`a}#&<-&F2l&x#EefGuw3U9t6Fb�n=d`xEu)+rlNXAYeWFZ@bA>)c(VlHo4Q zQ24`Rpx5AsWEkIc*`NO{meZOqT2SXx&L@@@)KLd zc1ib-@!p^um=t%%-37=gHvkyQm`vL~JMN$;qcUsY+2c5c?3lrfwpOW5V%R|w5B(mn zF-odPw1I1d-;LfB+JOr0%xd_)A!%@&pzvC;+OWDxKIK#`xv`bXGOOxqt2UHd-jN$| zf2g4L>$kzUMJ3PAF|W%OeJtUcPyfp|U=P0)FVWMnkKeXet~DErd{Au#%2nY|c=4K~ zTn9zRI2?i)@dgf#e%+~;{2`u!QkqTtMAp?t1|PmoMgP)n6PH=219A zcz%cZO@1kxjah?3KWaUyjdyfvB>O8IU{tH7)h(pmIdnaoCT2$WD`(Q%c0Zfpc5nsV zol2;-5$uf`oKG$hY+VLR&HApz>e$#&E~buvk72kr1AO5D**3ZI&g9<4p)WdTfoV&- zYAXmm;ubozo|hcG;a@*^TOAJ*@wP4QAAbzc`RQq)dq&M6wzB9T{tm#AU|v;~eE0QD z@@I5h3#Vs7B&h$D{Xvvdc=ln`)#spzGKIOc(yVm{_lz z(t0+lF&ZK)AcR{tRqLLqu59yo>p--ne=D;X^tO>ZsXYD@Of$Xbk#TV7xcgAw>482L zcVCz!NKD-cdsiZDsVR^DWv^?v;rlxgVBJ7qCx{O!dO5S<<`q7SV}4ot*)3q%jT;+o z3mxHs%+7e77+zz>Thew*n=L1vdeSTofVH*x{peSM8C{}`b2jo@a-4k2ta3|E1wr&k zEHV3>fL1Tz_ohpj_DOFcp=!DN=a{)yYO6$_Y5C0g1DC^>8sO@x`flY$8(!MQPodF9 zP}3En;ST3QoJKm2%SblI{PW^EYLp$3kvbw^U#pcTyAb@2r}G)T-^vzYoa$Wrkhgtm zWyw!-AaSCz|FFQqH7ANsk_)r$96VCE4uz~hYO6k)$_Lpl@2)fA8ZE)}!WIZF+08{i z@G$jk3op?|iV|c>6$!T-^KqWZuW;G6HSBo(V(hb6Mwv!ANqO7w3hzpIzZCuOr&sS& zb*Q#Bq5Mph?TTjJwJG0kalO3qL7wNFkPnu?D#EX0x_pxV#w1LaUsc9B9A5ODt8{LT zi*AVd=B-5(*&_60xW(=)`^lSf{;=Yj(bvCRmesww#fFl!;sm@&Zo2w_`ut0Fv`>oA zg3N>HdfkaN1v8JW%_uD`;`4a7m6s?_30DD|8+pZrEeA3heq8xb*useNYm6#&JrSDM zUv7!JA5lkx&xOkZ(V4FPR!&RtFrx)V?aw}MlVl(f3ENjU&{yd z<=Y>o3Gb&=n5S%s3D6~7wH08I)7`>d1o%O^+6}E0rH?*H3>nBDa?_S(Q3EVVYF)ROhCGN*kLkTQvia zUQ2t3o+0gsm#wk3Gu>rpNbf!U5=$Nk6fur(qp7b%U+K;}u}PIXPItw;*y)#NjomZ@ zTHWBAT5qJvmmNb-nnOYSB`1SOHXjhD1Z^{*7ZU2^1Sk@QaJyi7 zqO<2pz2A;hg5cHeMX{w!;hPE{N+$e}!hH)j6ffJY9pSiCEKBO8?4{Lu89s(b|Yvcoq}t5J1Kz~ZRvnnO3ftFK4o8tXWNQ|-?mPkZQ0WvUi#F7hS?Qj8nM;~ zDtBBm^l#vXBEXpOcwm+?InV8IHagb&AzGx2^h{1?8!o&+@~qLiqEAyReH-F9nh%O_ zYVMJ<3QTv)AY~v~X7#M%i;C^bc;=}pUy*ngVJ7OnKV;`K2hiAkNp=0FW}^?FoAvz+ z52_h#zoduIW^(A_$BUyMNlp;c7%Ey^1z$9 zalT>Y7r4i5HGgI+t-V$D#T%8RrDuoVurg*)gni)~9BwDe@$7vZ;w#9nur=&SaIz2U z&QQ!D1w038nFWDKs%_DatMCZOHx<5oLRWLk=f=;>Lx*{ZDa7!7?|n^0`#+f8$la6! zH_}Y*yC^H8uCs|yk)V6QTEn)fM#%BGPocowrUg1&crC;ht7n^GY~i~6#6 zE3do7R@J)OfRpRwA^?Lk4gULU2RH(EB5tO5D4r;kff7a5hFX$ms!L9h zpXn=xpCBr0fkf}qD^nzc#9gF{pS-s{J{v|CH#A_UX&cAz(d?B$Xk$wps z{m(||=u?OuHYVWvT#(6Sy0{Do2VV0Mb^OIot5<5??2ev|mHqekG?Cm03bPw1rges7 zn%{+<#E(0n{F)tg-ym7k+0(!LPPMl$_E(-d<}$D|Y07ai9Hb6mX6^tJIcvC_D$ z(%^m*ny=0#ch8NkTz6$s`%{!5Z%7?%W;xsKL4ES<3HdF*rFhcfxB)kZsI9R~cz?@Q z;FyEpZahrc&7^xfa{rP%G+cg*q1o4d?^AeT4&UK+eyX*F3QLKzWZcuU|NLC{qBuPy9g%xJQ9hES?q7d)3B9j5KpPX6HHs{rMo zOK~SW5O(qp1GG^cltiNZkx#||Tqh^I75cq`A?*4RKNi^nnW{11rG!wy)UjixbMS(&q@ zXQ(AYJ2KIC1bde6USSjSl zE)kq&i7L7=nK-0p*Wq00g-T$A6S&o?wGDc<$CG7 zY}Y(h(f-rlF3N?a_S~1gUK{&lVO{0UEnCHz-S@!NPcSqc32H&NYK=o+H06_={RxaY zJjh>%EQ3;|k%$nrMLtyB?cGK!#Kue9c9$t~9RZh>Z*}x+_@25H9}Uf`dx};0M5gV~ z3EMqLWoQkvwZH`<*)fYSR zwWuMpr521}Ee=3EyKOw6Hpt+58QhIWVzELlg+Af4#P_O&cXU&&Z2UxDoogkUzdCz! z5vFx`m8j;ZwoCSXoP7;NF!`5AEx7;9;Nqm&wQIGB%rcg?Px-0LlLveq0CJH($Rscq zMbo?gXkSviFCN&VUk|O4K+ELzHSnf=VC%#9$AAZ~et!7DH+lZw{^0Z6(o)15m$}RP z5UNj-+hTpw_AdP36Y{%eL1HZ>R}hXT2QA4G$t!fm0JnW=%5O%EqY)wZ{hcV`2Gl6D z&Jz;yDRx(psGfDYDKF@TS0$ZpaPVnHMu=40cHe^Y1i@Q_S|a7l$ajRMs9x7Yy()JN z$$q=~9t1S8v>F9Z*trpYZizhjkfsvt5t!fEBl^ToCT_wH`N^{Zq;f~nnw+k^PF~{m z-s`SRn!p83rfc88B#*q!3`+TAvQ<@fy4Qw@n$)=rBots$)V}MrRXfrwuCS#jjktTo z6Oaa?S9JN6C13Vq6EHZp+pr$zcAO=R86J%D6OJs`M1GQ>E`fE>nwZrQXG7e~m8^?C zY(~W3gA7(P{~Z1DjkVJkFgP(KlSE_sZ`2!ey@?*7^b#&#&){iLYTftGeUq&eb^6oS zl!{+QZ$Z`oe}b+V>U;tVbIWXoLJkQI%KwU~fM$LDk@lAXpo3oT_w^(;yw@N2LvLc~ z=#h))aOgXOXcT;^B|JG97H5ii*hmtlB1Q`}y=uoK-_~G#`5!51^@)CIB_qO)f-{+D z9DJ4g^))QM40 zr)KHmBm|KG0?7eFW&w|3RR=b)XNwAvGeeZX6h4>@0FLbBO_}@WLT%6&fN6Cg6v5Ag zzUbwtc~?C?b%zbcW`(_zF^N&%MQ~L6K*Gw_u;tA5PZhtSmzARAN!$cb2(yStetlD{ zKX^|D%RWUNMO1vQIBm^|>2oiW3kf%34K_WMkoE_1xG@Ru)}wO^z|BHk(k8iUF2Z^I z(nusxLn6uB7FK27yENQl;VO)n@AcLh$HeSYZ1uBh;Y#)$L4HWa3f~8y`sO-c(&tx@ zbT?&NXPN%nX}!m6{Blm04SJ7NQQY@$RRuLLx910vbwYam4O^NWcWqQ}K7SA+{!gEu z|1!cjt+K6_2pKCv$1o=6nMq#TfUHXp2W>*9Q20W)>u_8f8|a7*5EwP-RVbjP86A;- zp;ldlPCkI}4z2*E_Ozej*`Kb)GwG}+^v^V>0eFI(luy#hU}TnTPkR~5xlEMQ=fFgn zT1-V`Fzg;zRaH^^`*?n4+;hKy{^2URKgCC{q_%u1~jv@ zP6^s;yAggJh&+%)Q@QU}vafUe4|#me#|6YS5suFPT0qRufPtt9SItduHz2u*BbvNz z0!8wMV!2jNjJL<^&4AVdvTYd)Dr8kQMpd-M)$MrOVqSolKP>aj3d6>&08=RhPeEJ?Q=z^Fh#5TjZRD?yZkH-=Mtv%vb>)BnFMgb@jGQi+e{&oAvV@1O%W z^E&$~uC-$(YygRFb%O)(L*CgBmS9ns4O-db(lNk7z*B{Ur*SbbN}>Dm@%c+HiBriI zlyC!@+E(Mp&fu*QpQ2I2YJL>1f=cO0lLPA$8y45eH7Jn!H^{Nv#vJ0bgxj3J|Ri;u19o+Hqk5}t=~ZmhgwHU!6-0B^@?0;*z& z;D)%3@JT|vJEvNS`lgo8FQEzb$XkTduX0mkYLx?7NB=JE$*e;31LFo$cS#aIHDozX zT|(&kS>F{wz-$qiV+AY{jrn3Ahk-s}NYPv>lS2 zmWzwao4GK{Pg?Kr^!@U8%t$STPp^7icQmPl947ni9XPe5nf3mGptY$Rfv+=CJ%OeY z>Hg$uXV0mUJdswvcYS${OZtkL>-T;xGgKuLed7$~c)5GK!sbJG-cY!grU}G9kjD3_)j-NsUOy)pxxq2bka2MK6 zunHcvYncox2UXdljv_V^qgfX5GQx?Q3w%Si?UpFgrCVOJck92;ogBHDR??<277`#? zZymZzXQp@Ew^aj!0F6&81g1V%F2XuZ4QuUXE|VtC5xkpTFABnz!hsHF}q z@V`Z_gS@*F+*8sTKr7{63fMgpvqM=2z!BXI-XU_FC$7`*o}vb1`P~OCaFV|d-=RPf zZ=J-EW`wKt`Vx2jH5-Pu@t4|A=2D?}4(TOT-0h8{W6}Gvg%pF^C zfcGUZ&-2FW6E?h6qI*Z=slzzyLDmeGY52z$+`q79T;5nfUCq10kXAXWNcL0geXY~P z@?RzXvZL4arw=OO+GM}4=NHA4Glg#uz6bWY9_R8A&m;(5=LW2DbO3GBYnI-{FOgWWQnsCwJG;n*I+={;6 zBKTbkvh=S_c@jTvrHU4{&jFiZ0*#sEB9`JT96P!}1j}uxJT@HdsC1;CU!@1iYC}y* zI`w1WYIOp}9NCcwCe`9WO#%4G=Q1=rzObBG|G8j2D`$g9>Jx- zetv&ZvV%=j@)rrZ8mBC*lT&S-u@@O$z2jHOXi*XSpG2Mm=LQ0Ig)0$;5zSsr3eP_N zxGFztZw3dnmdJ4^BC!*32-klqosi!_0hvh%ABFxk{(aKuVxRuN@0*PiK|TQ5AyNfy@0)aefdH3{*sJ0w?>Pt*krsrGA^ zh#O}mVWl%!s<-3uD2=z$Gi+v4*f!+5#1bj*)p`(b+bYYfKX^@PnI2XW$M_Ey-pxEV z2{Fk$*1Ui|V&OIXaBg)wRjCbG36$o5NYXu}kL;nPhe|LCCmB1`W28!MT(>EAs^nh zy$};%b0UJ3C+O$h3HE7PKB85=5&SvWtZK_?Y*sadUdNX|oKBJV``ozE1D-I-z5qN% z!ZeuXJKqnbQ`qt;eK_h`VyH7|amTAMQ zb-IRwL@HIJHlJYxJSjPkkVY1R{H0 zzqt*W$?yW!QU1lUqQviwBzHQ6LJ2($Q>R32bG}>C|DO5&Q@SkZVn;3kE%C;cFRMhn z+kC$zoq4@fjc3vg4iG~i&|(k2ZHayh!5Yx%przxf2M#Z4iZm2{GdWfy;8cl!kCgYY zf}UWCt%|Y-4H6v$ld9|ga~%X}R?h=ti(7Yr?}~0F=J+1(n|Fo7S*GGpm=h&*${Ek= z=#fj0K%UmK$KgQ_E7%OGaNQz_vM&WY{luGR9wwY$ir)%M1uRj$42t3$X8Km!VxFHB zi1p{?!t7s9X`$GpBy&{iHvPhsBWvH3j`wGA120R&1fpSVMTz^8KW~%Kv9_N_xM#fB zmtMK7MMUhrfV++zO7OkpR&>+w$nd_?1SA?NI=fIP`gQJ*^Mha>?$H@J5g?}5&DnV27QUCO#bV2 z?B@f13MXEPeR4GtN(K}WDUctDI2rb;bzHWr>eJVtB_VZlC>kdghD`JSdUDC1B8msr z@7Qv@Wsff?eUJND-9?;7P{aMH#-eKZv}Y zBXZg;HIBF%%jxY&%m|CNKmE*6e>*PP;lfaK_s1tO(loR$!A72p_Yd_f4hEG=lj9b$ z;e#}Y@5(VjgPo-i(-BXV4dV~P^1|^tojqI@0T_Bgd+^F+IwsiEDZ%EclQhB+mm#Z) z0IVQ*589n{4J1NVKT;C7pGHwK6R=FNT z2^hPpEgT z-v}c4EK5GzW%PD$8)r0HRg`Rzw)UZQ{pf8g*^>3!^NeoeOpg@6owS#Qf=8m@x0)>f zGch8Rq4yQ7BLN#X(?GdBa}i_dLUbBH6@nT#K0z8nXZAb^=5*4P$88vc1R*>ePU6Sr z0e#D%EuoPVB8A4z$%+VhJ4QmtY&bLZ+v!n3Fhi^3w~hZ&CC5afx9f0Rfr&&4(Ns0y zcn}t*u>);qKN+b?Em*A=)#Hf;b5DCrf6ufP#2&V5+3)2OBq3lARu@LgP^>PkkJ9S? zEfkrVmlM51VHkZP3xKr4 z6O&0D7GVpp=%zR`6j=xQn5&8J?SPYvg2dPx6&ykN%_T>}=e|+jE!r<)pQ4IuU!9(a z0I+R3RCEBuayG84DmRxY8y8`NNjlL_G_xwC<3D!F*lRNxuPrz*hKnvMVOhQ%J35P- zB9LYZC06^cL~bU6$@^DuV`1SZiaL{BIhUQ^&S*~j!p;&uhR)2$&5P9%!z;{?`^!0V!MRBbPn3d<<(Y^Sv*;Y5v50&`*+W&)C4dpq@sBS)WGZ?d z#|fJ360jpCKd+{y)lm(eF+ zh3|y`R(greJP3#z$l2RDG_eBe?0h!#L}h6-d7V{h`>yRuyyQ4%chcdDFhMkN2jQOf z1_yGt*hC|Hbhu*d*0CwwJ4%LE3ep`R>R&I)Q1^sX%`gP{S!OsO$4*dfFmc{_1GRf*LxLHQB5I8mWVF9DOG>j(dmpl1nQtFq&2(=n=Hc z6SZ2DH)pqJMrF|xlya#Gab!@h4Ubz$MGsSZ7Sg4BNl!Ik1yz4wKN@-nr%_BYfXHV* z6;Z^p!!PY34}TKrc|{vd1&qI)E;~OCW?vYE_jXr1sfP+HDHINqE0NiLW4kH_7f}j-l4kv{2+sKCJ|Q%UDhAl9s+BRP_g0DiZdg+@?bYe{ ztB$|UiaLppeSjATZeobq$JKdwJmBFT_*HS`VXS(37~9!(I|*^Jlk%LIKrDO*X@~_c zu`~qZ_~QYk2egdV{vOu-RmGv{5m=n7wnk{JJ&|$YZ|980#)d=OSYP45aZCvnWT2G0 zQr-ERvgXH1&zzVp7Ola}h5IHu&H0TnqJDIKuGz!iKxsfROFV|G^J)7eQ%%jzoYW-5 zv7MwB4=59^>Y=z(cK4hD@_(#fGNRR)FBRkFx*wP^)o@OR*R^twu|pvzJt-V5j3`ew z2YfiCpQ>T8LyWOnNxly6{18Z zLjD&wCkn?i9T-Sa87>k7CoHwMe8#njWxLoOVZ6XYP>BWm#(K^A*8IWd< z4OkMJu@KpS5ItW8mLo2kapH_FtT9PZ#zu+Vh1cKwjK^g{Vc-lu+^+JKzNSb^026$& zLN2gSZsq@LZaXV4$aGvGUpt(<=#;x1n}Sit@Y1oNQ1|`LK|sR_Tj*qL(+xIIGcgA5 z%Ss!1oQ7rK?VW_?->P_WL=7G*1S!U7VVrYdBi66J6-sK#SE8imKkG{W!SQydZ~G5J z@RK+7n_7~ZL_%N>@jg1DoVnlSUjgIlmni?S#&=cXBWF-LZneOMn^lW#( zwNRGG$QjNJY^TsjZM$I*2-3LYH1=V_8~hrLqkiaD3Yd%B0_N}KJBeG0Xq%e;vT!p) zswIS>;!v)U^30ZKj}KUl8-SXc0JFWZOyM4mNepvRH{+?;>X0 z20X$3&vMv*myZ|;{b`+zTs~Ro@869nKuY0x*rL&lB0H&@{QefgSN9Y4fOZAi;$7aV z9pIkK-txLw;_@RA6Kv3TlEGk9zESLpM$%A+0s-Y_>d$W7R&%<^^Q2b1k}9Ow1M=mh zY%%C26cHabOTACQWfwMbgbj_}pE9CE>{EoRh0jJ{B4$D;yc~%7$Y%+mbLPH&9A6`U z_bJ{bCQp{(aw8nVzx9!=rKV4vW+b_cVhH$nU?x+3i={l{oKe9{7)X`!=Zwx&BSV3c zmvJNzQY6KnBnWj@!BxR@n4fp|98u!j>|fuycm2Ndj%A3`8F2-FIUFfHwRX9tvawXRP^NN6;QkmdPtIw5sRi`s7VkzhG~KMnC?LkM zq!tDJ$8r}SH*(~$RK5{7Z=#OnLHYoVpxtndWAD3o26w35&&3>kFEk5q;dsZaPto=UCm!-|HlrUMr zY?nC&9WA$~#z4C3jP>WmBo~9h*AK8MS<8|eFXYMOgU?^>B%92KAxq#Z%{~2R zQ@^U4q}zQi)BpB!@bKVG_ixFynLhMsAf7d|T88#tNBj;a%AuSQUs$>@=Jt+e(EjaZa axUII|>U$e*dvhB6w{Oou<)mFlFa8hv)qp7g literal 0 HcmV?d00001 diff --git a/moshi_src/LICENSE b/moshi_src/LICENSE new file mode 100644 index 0000000..31aa793 --- /dev/null +++ b/moshi_src/LICENSE @@ -0,0 +1,23 @@ +Permission is hereby granted, free of charge, to any +person obtaining a copy of this software and associated +documentation files (the "Software"), to deal in the +Software without restriction, including without +limitation the rights to use, copy, modify, merge, +publish, distribute, sublicense, and/or sell copies of +the Software, and to permit persons to whom the Software +is furnished to do so, subject to the following +conditions: + +The above copyright notice and this permission notice +shall be included in all copies or substantial portions +of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF +ANY KIND, EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED +TO THE WARRANTIES OF MERCHANTABILITY, FITNESS FOR A +PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT +SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY +CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION +OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, OUT OF OR +IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER +DEALINGS IN THE SOFTWARE. diff --git a/moshi_src/LICENSE.audiocraft b/moshi_src/LICENSE.audiocraft new file mode 100644 index 0000000..b93be90 --- /dev/null +++ b/moshi_src/LICENSE.audiocraft @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) Meta Platforms, Inc. and affiliates. + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/moshi_src/moshi/__init__.py b/moshi_src/moshi/__init__.py new file mode 100644 index 0000000..58c3afa --- /dev/null +++ b/moshi_src/moshi/__init__.py @@ -0,0 +1,19 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +""" +moshi is the inference codebase for Kyutai audio generation models. + +The code has been adapted from Audiocraft, see LICENSE.audiocraft + Copyright (c) Meta Platforms, Inc. and affiliates. +""" + +# flake8: noqa +from . import conditioners +from . import models +from . import modules +from . import quantization +from . import utils + +__version__ = "0.2.9a1" diff --git a/moshi_src/moshi/conditioners/__init__.py b/moshi_src/moshi/conditioners/__init__.py new file mode 100644 index 0000000..068cd41 --- /dev/null +++ b/moshi_src/moshi/conditioners/__init__.py @@ -0,0 +1,10 @@ +# flake8: noqa +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +""" +Modules to help doing generations under some fixed conditions. +""" + +from .base import (ConditionType, ConditionAttributes, ConditionFuser, ConditionProvider, + BaseConditioner, TensorCondition, ConditionTensors, dropout_all_conditions) diff --git a/moshi_src/moshi/conditioners/base.py b/moshi_src/moshi/conditioners/base.py new file mode 100644 index 0000000..7db3cf1 --- /dev/null +++ b/moshi_src/moshi/conditioners/base.py @@ -0,0 +1,432 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +# +# Adapted from +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. + +from collections import defaultdict +from dataclasses import dataclass, field +from itertools import chain +import logging +import typing as tp + +import torch +from torch import nn + +from ..modules.transformer import create_sin_embedding + + +logger = logging.getLogger(__name__) +TextCondition = tp.Optional[str] # a text condition can be a string or None (if doesn't exist) +ConditionTensors = dict[str, 'ConditionType'] + + +class ConditionType(tp.NamedTuple): + """Return type for a conditioner: both a condition tensor, and a mask indicating valid positions. + """ + condition: torch.Tensor + mask: torch.Tensor + + +@dataclass(frozen=True) +class TensorCondition: + """Looks quite similar to ConditionType, but represents the input to TensorConditioners. + `tensor` should be [B | 1, T, D], and `mask` should be `[B | 1, T]`. + """ + tensor: torch.Tensor + mask: torch.Tensor + + @staticmethod + def from_tensor(tensor: torch.Tensor): + B, T, _ = tensor.shape + mask = torch.ones(B, T, dtype=torch.bool, device=tensor.device) + return TensorCondition(tensor, mask) + + @staticmethod + def cat(conditions: tp.Sequence['TensorCondition']) -> 'TensorCondition': + assert conditions, "Cannot cat empty list." + ref_tensor = conditions[0].tensor + B, _, D = ref_tensor.shape + assert B == 1 + B = len(conditions) + T = max(condition.tensor.shape[1] for condition in conditions) + mask = torch.zeros(B, T, dtype=torch.bool, device=ref_tensor.device) + tensor = torch.zeros(B, T, D, dtype=ref_tensor.dtype, device=ref_tensor.device) + for b, condition in enumerate(conditions): + tensor[b, :condition.tensor.shape[1], :] = condition.tensor[0] + mask[b, :condition.mask.shape[1]] = condition.mask[0] + return TensorCondition(tensor, mask) + + +@dataclass +class ConditionAttributes: + """Standard class for representing the set of potential inputs to the conditioners. + Typically, `audiocraft.data.audio_dataset.SegmentInfo` will convert + to this class to make conditioning agnostic to the type of dataset. + + There are two kinds of conditionings: text (or None), or raw torch tensors (with a mask). + + """ + text: tp.Dict[str, tp.Optional[str]] = field(default_factory=dict) + tensor: tp.Dict[str, TensorCondition] = field(default_factory=dict) + + @property + def text_attributes(self) -> tp.Iterable[str]: + return self.text.keys() + + @property + def tensor_attributes(self) -> tp.Iterable[str]: + return self.text.keys() + + @staticmethod + def condition_types() -> tp.FrozenSet[str]: + return frozenset(["text", "tensor"]) + + def copy(self) -> 'ConditionAttributes': + return ConditionAttributes(dict(self.text), dict(self.tensor)) + + +Prepared = tp.TypeVar('Prepared') # represents the prepared condition input type. + + +class BaseConditioner(nn.Module, tp.Generic[Prepared]): + """Base model for all conditioner modules. + + Args: + dim (int): internal dim of the model. + output_dim (int): Output dim of the conditioner. + force_linear (bool, optional): Force linear projection even when `dim == output_dim`. + pad_empty (bool): if True, conditionings of 0 length will be padded to have length 1. + output_bias (bool): if True, the output projection will have a bias. + learn_padding (bool): if True, the padding value will be learnt, zero otherwise. + """ + + def __init__(self, dim: int, + output_dim: int, + device: tp.Union[torch.device, str], + force_linear: bool = True, + pad_empty: bool = True, + output_bias: bool = False, + learn_padding: bool = True): + super().__init__() + self.dim = dim + self.output_dim = output_dim + self.pad_empty = pad_empty + self.device = device + self.output_proj: nn.Module + if force_linear or dim != output_dim: + self.output_proj = nn.Linear(dim, output_dim, bias=output_bias, device=device) + assert not output_bias + else: + self.output_proj = nn.Identity() + self.learnt_padding: tp.Optional[torch.Tensor] + if learn_padding: + self.learnt_padding = nn.Parameter( + torch.randn(1, 1, output_dim, device=device), requires_grad=True) + self.learnt_padding.data *= 0.2 + else: + self.learnt_padding = None + + def prepare(self, *args, **kwargs) -> Prepared: + """Should be any part of the processing that will lead to a synchronization + point, e.g. BPE tokenization with transfer to the GPU. + + The returned value will be saved and return later when calling forward(). + """ + raise NotImplementedError() + + def _get_condition(self, inputs: Prepared) -> ConditionType: + """Gets input that should be used as conditioning (e.g, genre, description or a waveform). + Outputs a ConditionType, after the input data was embedded as a dense vector. + + Returns: + ConditionType: + - A tensor of size [B, T, dim] where B is the batch size, T is the length of the + output embedding and `dim` is the internal dimension of the embedding. + - And a mask indicating where the padding tokens, of shape `[B, T]`. + """ + raise NotImplementedError() + + def forward(self, inputs: Prepared) -> ConditionType: + cond, mask = self._get_condition(inputs) + B, T, C = cond.shape + if T == 0 and self.pad_empty: + cond = torch.zeros(B, T, C, device=cond.device, dtype=cond.dtype) + mask = torch.zeros_like(cond[..., 0], dtype=torch.bool) + + cond = self.output_proj(cond) + + maskf = mask.float()[..., None] + if self.learnt_padding is not None: + cond = cond * maskf + self.learnt_padding * (1 - maskf) + else: + cond = cond * maskf + return ConditionType(cond, mask) + + +class _BaseTextConditioner(BaseConditioner[Prepared]): + pass + + +class _BaseTensorConditioner(BaseConditioner[Prepared]): + pass + + +def dropout_tensor(condition: TensorCondition) -> TensorCondition: + """Utility function for nullifying a WavCondition object. + """ + return TensorCondition( + tensor=torch.zeros_like(condition.tensor), + mask=torch.zeros_like(condition.mask)) + + +def dropout_condition_(sample: ConditionAttributes, condition_type: str, condition: str) -> None: + """Utility function for nullifying an attribute inside a ConditionAttributes object. + Works in-place. + """ + valid_conditions = ConditionAttributes.condition_types() + if condition_type not in valid_conditions: + raise ValueError( + "dropout_condition got an unexpected condition type!" + f" expected one of {valid_conditions} but got '{condition_type}'") + + if condition not in getattr(sample, condition_type): + raise ValueError( + "dropout_condition received an unexpected condition!" + f" expected tensor={sample.tensor.keys()} and text={sample.text.keys()}" + f" but got '{condition}' of type '{condition_type}'!" + ) + + if condition_type == 'tensor': + tensor_condition = sample.tensor[condition] + sample.tensor[condition] = dropout_tensor(tensor_condition) + elif condition_type == 'text': + sample.text[condition] = None + else: + assert False + + +def dropout_all_conditions(attributes: tp.Sequence[ConditionAttributes]) -> list[ConditionAttributes]: + """ + Args: + attributes (list[ConditionAttributes]): All conditions attributes. + Returns: + list[ConditionAttributes]: Same with all conditions dropped. + """ + attributes = [attribute.copy() for attribute in attributes] + for condition_type in ConditionAttributes.condition_types(): + for attribute in attributes: + for condition in getattr(attribute, condition_type): + dropout_condition_(attribute, condition_type, condition) + return attributes + + +class ConditionProvider(nn.Module): + """Prepare and provide conditions given all the supported conditioners. + + Args: + conditioners (dict): Dictionary of conditioners. + device (torch.device or str, optional): Device for conditioners and output condition types. + """ + + def __init__(self, conditioners: tp.Dict[str, BaseConditioner], device: tp.Union[torch.device, str] = "cpu"): + super().__init__() + self.device = device + self.conditioners = nn.ModuleDict(conditioners).to(device) + + @property + def text_conditions(self): + return [k for k, v in self.conditioners.items() if isinstance(v, _BaseTextConditioner)] + + @property + def tensor_conditions(self): + return [k for k, v in self.conditioners.items() if isinstance(v, _BaseTensorConditioner)] + + def _collate_text(self, samples: tp.Sequence[ConditionAttributes]) -> tp.Dict[str, tp.List[tp.Optional[str]]]: + """Given a list of ConditionAttributes objects, compile a dictionary where the keys + are the attributes and the values are the aggregated input per attribute. + For example: + Input: + [ + ConditionAttributes(text={"genre": "Rock", "description": "A rock song with a guitar solo"}, wav=...), + ConditionAttributes(text={"genre": "Hip-hop", "description": "A hip-hop verse"}, wav=...), + ] + Output: + { + "genre": ["Rock", "Hip-hop"], + "description": ["A rock song with a guitar solo", "A hip-hop verse"] + } + + Args: + samples (list of ConditionAttributes): List of ConditionAttributes samples. + Returns: + dict[str, list[str, optional]]: A dictionary mapping an attribute name to text batch. + """ + out: tp.Dict[str, tp.List[tp.Optional[str]]] = defaultdict(list) + texts = [x.text for x in samples] + for text in texts: + for condition in self.text_conditions: + out[condition].append(text[condition]) + return out + + def _collate_tensors(self, samples: tp.Sequence[ConditionAttributes]) -> tp.Dict[str, TensorCondition]: + """For each tensor attribute, collate the tensor from individual batch items. + + Args: + samples (list of ConditionAttributes): List of ConditionAttributes samples. + Returns: + dict[str, TensorCondition]: A dictionary mapping an attribute name to tensor. + """ + per_attribute = defaultdict(list) + out: tp.Dict[str, TensorCondition] = {} + for sample in samples: + for attribute in self.tensor_conditions: + per_attribute[attribute].append(sample.tensor[attribute]) + + # stack all tensors to a single tensor + for attribute in self.tensor_conditions: + out[attribute] = TensorCondition.cat(per_attribute[attribute]) + + return out + + def prepare(self, inputs: tp.Sequence[ConditionAttributes]) -> tp.Dict[str, tp.Any]: + """Match attributes/tensors with existing conditioners in self, and call `prepare` for each one. + This should be called before starting any real GPU work to avoid synchronization points. + This will return a dict matching conditioner names to their arbitrary prepared representations. + + Args: + inputs (list[ConditionAttributes]): List of ConditionAttributes objects containing + text and tensors conditions. + """ + assert all([isinstance(x, ConditionAttributes) for x in inputs]), ( + "Got unexpected types input for conditioner! should be tp.List[ConditionAttributes]", + f" but types were {set([type(x) for x in inputs])}" + ) + + output = {} + text = self._collate_text(inputs) + tensors = self._collate_tensors(inputs) + + assert set(text.keys() | tensors.keys()).issubset(set(self.conditioners.keys())), ( + f"Got an unexpected attribute! Expected {self.conditioners.keys()}, ", + f"got {text.keys(), tensors.keys()}" + ) + + missing_inputs = set(self.conditioners.keys()) - (set(text.keys()) | set(tensors.keys())) + if missing_inputs: + raise RuntimeError(f"Some conditioners did not receive an input: {missing_inputs}") + for attribute, batch in chain(text.items(), tensors.items()): + conditioner = self.conditioners[attribute] + assert isinstance(conditioner, BaseConditioner) + output[attribute] = conditioner.prepare(batch) + return output + + def forward(self, prepared: tp.Dict[str, tp.Any]) -> tp.Dict[str, ConditionType]: + """Compute pairs of `(embedding, mask)` using the configured conditioners and the prepared representations. + The output is for example: + { + "genre": (torch.Tensor([B, 1, D_genre]), torch.Tensor([B, 1])), + "description": (torch.Tensor([B, T_desc, D_desc]), torch.Tensor([B, T_desc])), + ... + } + + Args: + prepared (dict): Dict of prepared representations as returned by `prepare()`. + """ + output = {} + for name, inputs in prepared.items(): + condition, mask = self.conditioners[name](inputs) + output[name] = ConditionType(condition, mask) + return output + + +class ConditionFuser(nn.Module): + """Condition fuser handles the logic to combine the different conditions + to the actual model input. + + Args: + fuse2cond (tp.Dict[str, str]): A dictionary that says how to fuse + each condition. For example: + { + "prepend": ["description"], + "sum": ["genre", "bpm"], + "cross": ["description"], + } + cross_attention_pos_emb (bool, optional): Use positional embeddings in cross attention. + cross_attention_pos_emb_scale (int): Scale for positional embeddings in cross attention if used. + """ + FUSING_METHODS = ["sum", "prepend", "cross"] + + def __init__(self, fuse2cond: tp.Dict[str, tp.List[str]], cross_attention_pos_emb: bool = False, + cross_attention_pos_emb_scale: float = 1.0): + super().__init__() + assert all( + [k in self.FUSING_METHODS for k in fuse2cond.keys()] + ), f"Got invalid fuse method, allowed methods: {self.FUSING_METHODS}" + self.cross_attention_pos_emb = cross_attention_pos_emb + self.cross_attention_pos_emb_scale = cross_attention_pos_emb_scale + self.fuse2cond: tp.Dict[str, tp.List[str]] = fuse2cond + self.cond2fuse: tp.Dict[str, str] = {} + for fuse_method, conditions in fuse2cond.items(): + for condition in conditions: + self.cond2fuse[condition] = fuse_method + if fuse_method not in ['cross', 'sum']: + raise RuntimeError("only `sum` and `cross` conditionings are supported " + f"for now, got {fuse_method}.") + + @property + def has_conditions(self) -> bool: + return bool(self.cond2fuse) + + @property + def has_prepend(self) -> bool: + """Is there a conditioning that needs to be prepending to the Transformer sequence.""" + return bool(self.fuse2cond['prepend']) + + def get_cross(self, conditions: ConditionTensors) -> torch.Tensor | None: + """Return the tensor to be provided for the cross attention.""" + cross = None + for name in self.fuse2cond['cross']: + cond, _ = conditions[name] + if cross is None: + cross = cond + else: + cross = torch.cat([cross, cond], dim=1) + + if self.cross_attention_pos_emb and cross is not None: + positions = torch.arange( + cross.shape[1], + device=cross.device + ).view(1, -1, 1) + pos_emb = create_sin_embedding(positions, cross.shape[-1]).to(cross.dtype) + cross = cross + self.cross_attention_pos_emb_scale * pos_emb + return cross + + def get_sum(self, conditions: ConditionTensors) -> torch.Tensor | None: + """Return the tensor to be provided as an extra sum offset shared for each step.""" + sum = None + for name in self.fuse2cond['sum']: + cond, _ = conditions[name] + assert cond.shape[1] == 1, cond.shape + if sum is None: + sum = cond + else: + sum = sum + cond + return sum + + def get_prepend(self, conditions: ConditionTensors) -> torch.Tensor | None: + """Return the tensor to be prepended to the transformer.""" + prepend = None + for name in self.fuse2cond['prepend']: + cond, _ = conditions[name] + if prepend is None: + prepend = cond + else: + prepend = torch.cat([cond, prepend], dim=1) + if prepend is not None: + sum = self.get_sum(conditions) + if sum is not None: + prepend = prepend + sum + return prepend diff --git a/moshi_src/moshi/conditioners/tensors.py b/moshi_src/moshi/conditioners/tensors.py new file mode 100644 index 0000000..40c0cc0 --- /dev/null +++ b/moshi_src/moshi/conditioners/tensors.py @@ -0,0 +1,16 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +from .base import _BaseTensorConditioner, TensorCondition, ConditionType + + +class TensorConditioner(_BaseTensorConditioner[TensorCondition]): + """Does basically nothing. + """ + + def prepare(self, tensor: TensorCondition) -> TensorCondition: + device = next(iter(self.parameters())).device + return TensorCondition(tensor.tensor.to(device=device), tensor.mask.to(device=device)) + + def _get_condition(self, inputs: TensorCondition) -> ConditionType: + return ConditionType(inputs.tensor, inputs.mask) diff --git a/moshi_src/moshi/conditioners/text.py b/moshi_src/moshi/conditioners/text.py new file mode 100644 index 0000000..10dc5ab --- /dev/null +++ b/moshi_src/moshi/conditioners/text.py @@ -0,0 +1,134 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +import hashlib +import logging +import typing as tp + +import torch +from torch import nn + + +from .base import _BaseTextConditioner, ConditionType + + +logger = logging.getLogger(__name__) + + +def length_to_mask(lengths: torch.Tensor, max_len: tp.Optional[int] = None) -> torch.Tensor: + """Utility function to convert a tensor of sequence lengths to a mask (useful when working on padded sequences). + For example: [3, 5] => [[1, 1, 1, 0, 0], [1, 1, 1, 1, 1]] + + Args: + lengths (torch.Tensor): tensor with lengths + max_len (int): can set the max length manually. Defaults to None. + Returns: + torch.Tensor: mask with 0s where there is pad tokens else 1s + """ + assert len(lengths.shape) == 1, "Length shape should be 1 dimensional." + final_length = lengths.max().item() if not max_len else max_len + final_length = max(final_length, 1) # if all seqs are of len zero we don't want a zero-size tensor + return torch.arange(final_length, device=lengths.device)[None, :] < lengths[:, None] + + +def hash_trick(word: str, vocab_size: int) -> int: + """Hash trick to pair each word with an index + + Args: + word (str): word we wish to convert to an index + vocab_size (int): size of the vocabulary + Returns: + int: index of the word in the embedding LUT + """ + hash = int(hashlib.sha256(word.encode("utf-8")).hexdigest(), 16) + return hash % vocab_size + + +class TokenizedText(tp.NamedTuple): + tokens: torch.Tensor # should be long tensor. + mask: torch.Tensor # should be bool tensor. + + +class TextConditioner(_BaseTextConditioner[TokenizedText]): + ... + + +class Tokenizer: + """Base tokenizer implementation + """ + def __call__(self, texts: tp.List[tp.Optional[str]]) -> TokenizedText: + raise NotImplementedError() + + +class NoopTokenizer(Tokenizer): + """This tokenizer should be used for global conditioners such as: artist, genre, key, etc. + The difference between this and WhiteSpaceTokenizer is that NoopTokenizer does not split + strings, so "Jeff Buckley" will get it's own index. Whereas WhiteSpaceTokenizer will + split it to ["Jeff", "Buckley"] and return an index per word. + + For example: + ["Queen", "ABBA", "Jeff Buckley"] => [43, 55, 101] + ["Metal", "Rock", "Classical"] => [0, 223, 51] + + When all possible values are known, one can use `possible_values` to provide the list + of possible tokens. If a token doesn't exist, `pad_idx` will be used instead. + """ + def __init__(self, n_bins: int, possible_values: list[str] | None = None): + self.n_bins = n_bins + self.pad_idx = n_bins + if possible_values is None: + self.possible_values = None + else: + self.possible_values = {value: idx for idx, value in enumerate(possible_values)} + assert n_bins >= len(possible_values) + + def __call__(self, texts: tp.List[tp.Optional[str]]) -> TokenizedText: + output, lengths = [], [] + for text in texts: + # if current sample doesn't have a certain attribute, replace with pad token + if text is None: + output.append(self.pad_idx) + lengths.append(0) + else: + if self.possible_values is None: + output.append(hash_trick(text, self.n_bins)) + else: + if text not in self.possible_values: + raise ValueError(f"'{text}' is not in possible_values {self.possible_values}") + output.append(self.possible_values[text]) + lengths.append(1) + + tokens = torch.tensor(output).int()[:, None] + mask = length_to_mask(torch.tensor(lengths)) + return TokenizedText(tokens, mask) + + +class LUTConditioner(TextConditioner): + """Lookup table TextConditioner. + + Args: + n_bins (int): Number of bins. + dim (int): Hidden dim of the model (text-encoder/LUT). + output_dim (int): Output dim of the conditioner. + pad_idx (int, optional): Index for padding token. Defaults to 0. + """ + def __init__(self, n_bins: int, tokenizer: str, possible_values: list[str] | None = None, + init_scale: float = 1., **kwargs): + super().__init__(**kwargs) + self.embed = nn.Embedding(n_bins + 1, self.dim) # n_bins + 1 for padding. + self.embed.weight.data *= init_scale + if tokenizer == 'noop': + self.tokenizer = NoopTokenizer(n_bins, possible_values) + else: + raise ValueError(f"unrecognized tokenizer `{tokenizer}`.") + + def prepare(self, x: tp.List[tp.Optional[str]]) -> TokenizedText: + device = self.embed.weight.device + tokens, mask = self.tokenizer(x) + tokens, mask = tokens.to(device), mask.to(device) + return TokenizedText(tokens.to(device), mask.to(device)) + + def _get_condition(self, inputs: TokenizedText) -> ConditionType: + tokens, mask = inputs + embeds = self.embed(tokens) + return ConditionType(embeds, mask) diff --git a/moshi_src/moshi/models/__init__.py b/moshi_src/moshi/models/__init__.py new file mode 100644 index 0000000..85b1fb2 --- /dev/null +++ b/moshi_src/moshi/models/__init__.py @@ -0,0 +1,14 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +""" +Models for the compression model Moshi, +""" + +# flake8: noqa +from .compression import ( + CompressionModel, + MimiModel, +) +from .lm import LMModel, LMGen +from .loaders import get_mimi, get_moshi_lm diff --git a/moshi_src/moshi/models/compression.py b/moshi_src/moshi/models/compression.py new file mode 100644 index 0000000..c1ae0d3 --- /dev/null +++ b/moshi_src/moshi/models/compression.py @@ -0,0 +1,488 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Part of this file is adapted from encodec.py in https://github.com/facebookresearch/audiocraft +# released under the following license. +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +"""Compression models or wrapper around existing models. In particular, provides the implementation +for Mimi. Also defines the main interface that a model must follow to be usable as an audio tokenizer. +""" + +from abc import abstractmethod +from dataclasses import dataclass +import logging +import typing as tp + +import torch +from torch import nn + + +from ..quantization import ( + QuantizedResult, + BaseQuantizer, + SplitResidualVectorQuantizer, + ResidualVectorQuantizer, +) +from ..modules.conv import pad_for_conv1d +from ..modules.resample import ConvDownsample1d, ConvTrUpsample1d +from ..modules.streaming import StreamingModule, State, StateT +from ..utils.compile import CUDAGraphed + + +logger = logging.getLogger() + + +class CompressionModel(StreamingModule[StateT]): + """Base API for all compression model that aim at being used as audio tokenizers + with a language model. + """ + + @abstractmethod + def forward(self, x: torch.Tensor) -> QuantizedResult: ... + + @abstractmethod + def encode(self, x: torch.Tensor) -> torch.Tensor: + """See `MimiModel.encode`.""" + ... + + @abstractmethod + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """See `MimiModel.decode`.""" + ... + + @abstractmethod + def decode_latent(self, codes: torch.Tensor) -> torch.Tensor: + """Decode from the discrete codes to continuous latent space.""" + ... + + @property + @abstractmethod + def channels(self) -> int: ... + + @property + @abstractmethod + def frame_size(self) -> int: ... + + @property + @abstractmethod + def frame_rate(self) -> float: ... + + @property + @abstractmethod + def sample_rate(self) -> int: ... + + @property + @abstractmethod + def cardinality(self) -> int: ... + + @property + @abstractmethod + def num_codebooks(self) -> int: ... + + @property + @abstractmethod + def total_codebooks(self) -> int: ... + + @abstractmethod + def set_num_codebooks(self, n: int): + """Set the active number of codebooks used by the quantizer.""" + ... + + +@dataclass +class _MimiState(State): + graphed_tr_enc: CUDAGraphed | None + graphed_tr_dec: CUDAGraphed | None + graphed_encoder: CUDAGraphed + graphed_decoder: CUDAGraphed + + +class MimiModel(CompressionModel[_MimiState]): + """Mimi model operating on the raw waveform. + + Args: + encoder (nn.Module): Encoder network. + decoder (nn.Module): Decoder network. + quantizer (qt.BaseQuantizer): Quantizer network. + frame_rate (float): Final frame rate of the quantized representatiopn. + encoder_frame_rate (float): frame rate of the encoder model. Note that if `frame_rate != encopder_frame_rate`, + the latent will be resampled linearly to match the desired `frame_rate` before and after quantization. + sample_rate (int): Audio sample rate. + channels (int): Number of audio channels. + causal (bool): Whether to use a causal version of the model. + encoder_transformer (nn.Module or None): optional transformer for the encoder. + decoder_transformer (nn.Module or None): optional transformer for the decoder. + resample_method (str): method to use for resampling the latent space before the quantizer. + upsample_channel_wise_bug (bool): controls whether the upsampling is channel wise. + Defaults to true to reproduce bug in original implementation. + freeze_encoder: whether to freeze the encoder weights. + freeze_quantizer: whether to freeze the quantizer weights. + freeze_quantizer_level: If positive, freeze the quantizer up to this level. + """ + + def __init__( + self, + encoder: nn.Module, + decoder: nn.Module, + quantizer: BaseQuantizer, + frame_rate: float, + encoder_frame_rate: float, + sample_rate: int, + channels: int, + causal: bool = False, + encoder_transformer: tp.Optional[nn.Module] = None, + decoder_transformer: tp.Optional[nn.Module] = None, + resample_method: str = "interpolate", + upsample_channel_wise_bug: bool = True, + freeze_encoder: bool = False, + freeze_quantizer: bool = False, + freeze_quantizer_level: int = -1, + ): + super().__init__() + self.encoder = encoder + self.decoder = decoder + self.encoder_transformer = encoder_transformer + self.decoder_transformer = decoder_transformer + self.quantizer = quantizer + self._frame_rate = frame_rate + self._sample_rate = sample_rate + self._channels = channels + self.encoder_frame_rate = encoder_frame_rate + + if freeze_encoder: + for p in self.encoder.parameters(): + p.requires_grad = False + if self.encoder_transformer is not None: + for p in self.encoder_transformer.parameters(): + p.requires_grad = False + for name, p in self.quantizer.named_parameters(): + if name.endswith("input_proj.weight"): + p.requires_grad = False + if freeze_quantizer: + self.quantizer.ema_frozen_(True) + self.freeze_quantizer = freeze_quantizer + self.freeze_quantizer_level = ( + freeze_quantizer_level + if freeze_quantizer_level > 0 + else self.quantizer.num_codebooks + ) + + # We will need the dimension for the resampling. In general the encoder will be a SeanetEncoder + # which exposes a `dimension` attribute. + dimension = encoder.dimension + assert isinstance( + dimension, int + ), f"Dimension should be int, got {dimension} of type {type(dimension)}." + self.dimension = dimension + + assert resample_method in [ + "interpolate", + "conv", + "avg_pool", + ], f"Invalid resample_method {resample_method}" + self.resample_method = resample_method + if encoder_frame_rate != frame_rate: + assert not ( + causal and resample_method == "interpolate" + ), "Cannot interpolate with causal model." + if resample_method in ["conv", "avg_pool"]: + assert ( + self.encoder_frame_rate > self.frame_rate + ), "Cannot upsample with conv." + downsample_stride = self.encoder_frame_rate / self.frame_rate + assert downsample_stride == int( + downsample_stride + ), f"Only integer strides are supported, got {downsample_stride}" + learnt = resample_method == "conv" + self.downsample = ConvDownsample1d( + int(downsample_stride), + dimension=dimension, + learnt=learnt, + causal=causal, + ) + if freeze_encoder: + for p in self.downsample.parameters(): + p.requires_grad = False + self.upsample = ConvTrUpsample1d( + int(downsample_stride), + dimension=dimension, + learnt=learnt, + causal=causal, + channel_wise=upsample_channel_wise_bug, + ) + + def _init_streaming_state(self, batch_size: int) -> _MimiState: + device = next(self.parameters()).device + disable = device.type != 'cuda' + graphed_tr_dec = None + graphed_tr_enc = None + if self.encoder_transformer is not None: + graphed_tr_enc = CUDAGraphed(self.encoder_transformer, disable=disable) + if self.decoder_transformer is not None: + graphed_tr_dec = CUDAGraphed(self.decoder_transformer, disable=disable) + graphed_encoder = CUDAGraphed(self.encoder, disable=disable) + graphed_decoder = CUDAGraphed(self.decoder, disable=disable) + return _MimiState(batch_size, device, graphed_tr_enc, graphed_tr_dec, graphed_encoder, graphed_decoder) + + @property + def channels(self) -> int: + return self._channels + + @property + def frame_rate(self) -> float: + return self._frame_rate + + @property + def sample_rate(self) -> int: + return self._sample_rate + + @property + def frame_size(self) -> int: + return int(self.sample_rate / self.frame_rate) + + @property + def total_codebooks(self): + """Total number of quantizer codebooks available.""" + return self.quantizer.total_codebooks + + @property + def num_codebooks(self): + """Active number of codebooks used by the quantizer.""" + return self.quantizer.num_codebooks + + def set_num_codebooks(self, n: int): + """Set the active number of codebooks used by the quantizer.""" + self.quantizer.set_num_codebooks(n) + + @property + def cardinality(self): + """Cardinality of each codebook.""" + return self.quantizer.cardinality + + def _to_framerate(self, x: torch.Tensor): + # Convert from the encoder frame rate to the overall framerate. + _, _, length = x.shape + frame_rate = self.encoder_frame_rate + new_frame_rate = self.frame_rate + if frame_rate == new_frame_rate: + return x + if self.resample_method == "interpolate": + target_length = int(length * new_frame_rate / frame_rate) + return nn.functional.interpolate(x, size=target_length, mode="linear") + else: + return self.downsample(x) + + def _to_encoder_framerate(self, x: torch.Tensor): + # Convert from overall framerate to the encoder frame rate. + _, _, length = x.shape + frame_rate = self.encoder_frame_rate + new_frame_rate = self.frame_rate + if frame_rate == new_frame_rate: + return x + if self.resample_method == "interpolate": + target_length = int(length * new_frame_rate / frame_rate) + return nn.functional.interpolate(x, size=target_length, mode="linear") + else: + return self.upsample(x) + + def forward(self, x: torch.Tensor) -> QuantizedResult: + assert x.dim() == 3 + length = x.shape[-1] + extra_metrics: tp.Dict[str, torch.Tensor] = {} + + if self.freeze_quantizer: + if isinstance(self.quantizer, SplitResidualVectorQuantizer): + self.quantizer.rvq_first.eval() + for i in range( + self.freeze_quantizer_level - self.quantizer.n_q_semantic + ): + self.quantizer.rvq_rest.vq.layers[i].eval() + elif isinstance(self.quantizer, ResidualVectorQuantizer): + for i in range(self.freeze_quantizer_level): + self.quantizer.vq.layers[i].eval() + else: + raise ValueError(f"Unsupported quantizer type {type(self.quantizer)}") + + emb = self.encoder(x) + if self.encoder_transformer is not None: + (emb,) = self.encoder_transformer(emb) + emb = self._to_framerate(emb) + expected_length = self.frame_rate * length / self.sample_rate + # Checking that we have the proper length given the advertised frame rate. + assert abs(emb.shape[-1] - expected_length) < 1, ( + emb.shape[-1], + expected_length, + ) + + q_res = self.quantizer(emb, self.frame_rate) + emb = q_res.x + emb = self._to_encoder_framerate(emb) + if self.decoder_transformer is not None: + (emb,) = self.decoder_transformer(emb) + + out = self.decoder(emb) + + # remove extra padding added by the encoder and decoder + assert out.shape[-1] >= length, (out.shape[-1], length) + out = out[..., :length] + + q_res.x = out + q_res.metrics.update(extra_metrics) + return q_res + + def _encode_to_unquantized_latent(self, x: torch.Tensor) -> torch.Tensor: + """Projects a batch of waveforms to unquantized latent space. + + Args: + x (torch.Tensor): Float tensor of shape [B, C, T]. + + Returns: + Unquantized embeddings. + """ + assert ( + x.dim() == 3 + ), f"CompressionModel._encode_to_unquantized_latent expects audio of shape [B, C, T] but got {x.shape}" + + state = self._streaming_state + frame_size = self.frame_size + + if state is None: + # The underlying convolutions no longer accept partial inputs, + # `x` needs to be exactly a multiple of the frame size, + # reproducing the previous padding behavior here. + x = pad_for_conv1d(x, frame_size, frame_size) + emb = self.encoder(x) + else: + if x.shape[-1] % frame_size != 0 or x.shape[-1] == 0: + raise RuntimeError( + f"Invalid input x of length {x.shape[-1]}. The length must be " + f"a positive multiple of the frame size {frame_size}. " + "You are responsible for buffering accordingly before feeding audio to Mimi.") + emb = state.graphed_encoder(x).clone() + if self.encoder_transformer is not None: + if state is None: + (emb,) = self.encoder_transformer(emb) + else: + assert state.graphed_tr_enc is not None + (emb,) = state.graphed_tr_enc(emb) + emb = self._to_framerate(emb) + return emb + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """Encode the given input tensor to quantized representation. + + Args: + x (torch.Tensor): Float tensor of shape [B, C, T] + + Returns: + codes (torch.Tensor): an int tensor of shape [B, K, T] + with K the number of codebooks used and T the timestep. + """ + emb = self._encode_to_unquantized_latent(x) + codes = self.quantizer.encode(emb) + return codes + + def encode_to_latent(self, x: torch.Tensor, quantize: bool = True) -> torch.Tensor: + """Projects a batch of waveforms to latent space. + + Args: + x (torch.Tensor): Float tensor of shape [B, C, T]. + + Returns: + Embeddings, either quantized or not. + """ + emb = self._encode_to_unquantized_latent(x) + if not quantize: + return emb + else: + codes = self.quantizer.encode(emb) + return self.decode_latent(codes) + + def decode(self, codes: torch.Tensor): + """Decode the given codes to a reconstructed representation. + + Args: + codes (torch.Tensor): Int tensor of shape [B, K, T] + + Returns: + out (torch.Tensor): Float tensor of shape [B, C, T], the reconstructed audio. + """ + state = self._streaming_state + emb = self.decode_latent(codes) + emb = self._to_encoder_framerate(emb) + if self.decoder_transformer is not None: + if state is None: + (emb,) = self.decoder_transformer(emb) + else: + assert state.graphed_tr_dec is not None + (emb,) = state.graphed_tr_dec(emb) + if state is None: + out = self.decoder(emb) + else: + out = state.graphed_decoder(emb).clone() + # out contains extra padding added by the encoder and decoder + return out + + def decode_latent(self, codes: torch.Tensor) -> torch.Tensor: + """Decode from the discrete codes to continuous latent space.""" + return self.quantizer.decode(codes) + + +class WrapperCompressionModel(CompressionModel[State]): + """Base API for CompressionModel wrappers that do not depend on external frameworks.""" + + def __init__(self, model: CompressionModel): + super().__init__() + self.model = model + + def forward(self, x: torch.Tensor) -> QuantizedResult: + return self.model.forward(x) + + def encode(self, x: torch.Tensor) -> torch.Tensor: + return self.model.encode(x) + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + return self.model.decode(codes) + + def decode_latent(self, codes: torch.Tensor) -> torch.Tensor: + return self.model.decode_latent(codes) + + def set_num_codebooks(self, n: int): + self.model.set_num_codebooks(n) + + @property + def quantizer(self): + return self.model.quantizer + + @property + def channels(self) -> int: + return self.model.channels + + @property + def frame_rate(self) -> float: + return self.model.frame_rate + + @property + def sample_rate(self) -> int: + return self.model.sample_rate + + @property + def frame_size(self) -> int: + return self.model.frame_size + + @property + def cardinality(self) -> int: + return self.model.cardinality + + @property + def num_codebooks(self) -> int: + return self.model.num_codebooks + + @property + def total_codebooks(self) -> int: + return self.model.total_codebooks diff --git a/moshi_src/moshi/models/lm.py b/moshi_src/moshi/models/lm.py new file mode 100644 index 0000000..2fc6f1c --- /dev/null +++ b/moshi_src/moshi/models/lm.py @@ -0,0 +1,837 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +from contextlib import ExitStack +from dataclasses import dataclass, field +from functools import partial +import logging +import typing as tp +import torch +from torch import nn +from ..conditioners import ConditionProvider, ConditionFuser, ConditionTensors +from ..utils.sampling import sample_token +from ..utils.compile import CUDAGraphed +from ..utils.quantize import replace_linear_with_qlinear +from ..modules.streaming import StreamingContainer, StreamingModule, State +from ..modules.transformer import StreamingTransformer, create_norm_fn +from .lm_utils import (_delay_sequence, + _undelay_sequence, + _init_layer, + ScaledEmbedding) + + +logger = logging.getLogger(__name__) + + +def scatter_with_mask_(tensor: torch.Tensor, dim: int, + index: torch.Tensor, value: torch.Tensor, mask: torch.Tensor) -> None: + """Scatter but skipping the updates that are masked.""" + old_value = tensor.gather(dim, index) + value = torch.where(mask, value, old_value) + tensor.scatter_(dim, index, value) + + +@dataclass +class LMOutput: + # The logits are already re-aligned with the input codes + # hence no extra shift is required, e.g. when computing CE + logits: torch.Tensor # [B, K, T, card] + mask: torch.Tensor # [B, K, T] + text_logits: torch.Tensor # [B, 1, T, text_card] + text_mask: torch.Tensor # [B, 1, T] + + +class LMModel(StreamingContainer): + """Transformer-based language model on multiple streams of codes. + + Args: + n_q (int): Number of parallel streams to model as input. + dep_q (int): Number of parallel streams to model in the depformer. + card (int): Cardinality, vocabulary size. + text_card (int): Cardinality of the text vocabulary. + text_card_out (int or None): Cardinality of output text, if different from the input. + demux_second_text_stream: (bool): Whether two text streams are muxed together with a cartesian product. + dim (int): Dimension of the transformer encoder. + num_heads (int): Number of heads for the transformer encoder. + hidden_scale (int): Scale for hidden feed forward dimension of the transformer encoder. + norm (str): Normalization method. + norm_emb (bool): Whether to normalize embeddings. + bias_proj (bool): Use bias for output projections. + depformer_*: params used for the Depformer Transformer, all the other will be shared. + depformer_multi_linear (bool): if True, uses one linear layer per codebook to project the + output of the main transformer to the Depformer latent space. + depformer_dim_feedforward (int| list[int]| None): If None, defaults to hidden_scale * depformer_dim. + depformer_weights_per_step_schedule (list[int] | None): mapping `CODEBOOK_INDEX -> WEIGHT_INDEX`, allowing + depformer_low_rank_embeddings (int | None): if provided, uses low rank embeddings, with a linear + existing_text_padding_id (int): token to use for the padding. + same_initial (bool): if True, uses the same initial tokens for both text and audio mode. + **kwargs: Additional parameters for the transformer encoder. + """ + + def __init__( + self, + delays: tp.List[int] = [0], + n_q: int = 8, + dep_q: int = 8, + card: int = 1024, + text_card: int = 32000, + text_card_out: int | None = None, + demux_second_text_stream: bool = False, + dim: int = 128, + num_heads: int = 8, + hidden_scale: int = 4, + norm: str = "layer_norm", + norm_emb: bool = False, + bias_proj: bool = False, + depformer_dim: int = 256, + depformer_dim_feedforward: int | list[int] | None = None, + depformer_multi_linear: bool = False, + depformer_weights_per_step: bool = False, + depformer_weights_per_step_schedule: list[int] | None = None, + depformer_low_rank_embeddings: int | None = None, + depformer_pos_emb: str = "sin", + existing_text_padding_id: int = 3, + existing_text_end_padding_id: int = 0, + extra_heads_num_heads: int = 0, + extra_heads_dim: int = 6, + context: tp.Optional[int] = None, + causal: bool = True, + condition_provider: tp.Optional[ConditionProvider] = None, + fuser: tp.Optional[ConditionFuser] = None, + quantize: bool = False, + device=None, + dtype=None, + gradient_checkpointing: bool = False, + **kwargs, + ): + super().__init__() + self.n_q = n_q + self.dep_q = dep_q + self.card = card + self.text_card = text_card + text_card_out = text_card if text_card_out is None else text_card_out + assert len(delays) == self.num_codebooks, f"expected {self.num_codebooks} delays, got {len(delays)}." + self.delays = delays + self.dim = dim + self.existing_text_padding_id = existing_text_padding_id + self.existing_text_end_padding_id = existing_text_end_padding_id + self.context = context + self.depformer_weights_per_step_schedule = depformer_weights_per_step_schedule + if depformer_weights_per_step_schedule is not None: + assert len(depformer_weights_per_step_schedule) == dep_q + EmbeddingFactory = partial( + ScaledEmbedding, + norm=norm_emb, + device=device, + dtype=dtype, + zero_idx=self.zero_token_id, + ) + self.emb = nn.ModuleList( + [EmbeddingFactory(self.card + 1, dim) for _ in range(n_q)] + ) + # Unlike for audio, here we authorize the model to output the special token. + self.text_emb = EmbeddingFactory(text_card + 1, dim, demux_second_stream=demux_second_text_stream) + + self.text_linear = nn.Linear(dim, text_card_out, bias=bias_proj) + depformer_prefix = "depformer_" + main_kwargs = { + k: v for k, v in kwargs.items() if not k.startswith(depformer_prefix) + } + self.transformer = StreamingTransformer( + d_model=dim, + num_heads=num_heads, + dim_feedforward=int(hidden_scale * dim), + norm=norm, + device=device, + dtype=dtype, + quantize=quantize, + context=context, + causal=causal, + checkpointing=gradient_checkpointing, + **main_kwargs, + ) + self.out_norm = create_norm_fn(norm, dim) + self.depformer_multi_linear = depformer_multi_linear + kwargs_dep = main_kwargs.copy() + kwargs_dep.update( + { + k.removeprefix(depformer_prefix): v + for k, v in kwargs.items() + if k.startswith(depformer_prefix) + } + ) + kwargs_dep["positional_embedding"] = depformer_pos_emb + kwargs_dep["context"] = None + kwargs_dep["cross_attention"] = False + if depformer_weights_per_step: + kwargs_dep["weights_per_step"] = dep_q + if depformer_multi_linear: + # One linear layer per codebook to project different informations from the main model. + num_in = dep_q + if depformer_weights_per_step_schedule: + num_in = max(depformer_weights_per_step_schedule) + 1 + self.depformer_in = nn.ModuleList( + [nn.Linear(dim, depformer_dim, bias=False) for _ in range(num_in)] + ) + else: + self.depformer_in = nn.ModuleList( + [nn.Linear(dim, depformer_dim, bias=False)] + ) + EmbeddingFactory = partial(EmbeddingFactory, low_rank=depformer_low_rank_embeddings) + if dep_q > 0: + # Only using up to dep_q - 1 because the last codebook is never an input to Depformer. + self.depformer_emb = nn.ModuleList( + [EmbeddingFactory(self.card + 1, depformer_dim) for _ in range(dep_q - 1)] + ) + self.depformer_text_emb = EmbeddingFactory( + text_card + 1, + depformer_dim, + demux_second_stream=demux_second_text_stream, + ) + if depformer_dim_feedforward is None: + depformer_dim_feedforward = int(hidden_scale * depformer_dim) + self.depformer = StreamingTransformer( + d_model=depformer_dim, + dim_feedforward=depformer_dim_feedforward, + norm=norm, + weights_per_step_schedule=depformer_weights_per_step_schedule, + causal=causal, + quantize=quantize, + checkpointing=gradient_checkpointing, + device=device, + dtype=dtype, + **kwargs_dep, + ) + # Depformer follow its own cycle of streaming entirely contained in one time step + # and should not follow the streaming of the steps dimensions. + self.depformer.set_streaming_detached(True) + else: # No-Depformer --- e.g., an ASR model + self.depformer_emb = None + self.depformer_text_emb = None + self.depformer = None + + self.extra_heads = nn.ModuleList( + [nn.Linear(dim, extra_heads_dim, bias=False) for _ in range(extra_heads_num_heads)] + ) + + dim = depformer_dim # we will directly apply the next linears to the output of the Depformer. + + self.linears = nn.ModuleList( + [nn.Linear(dim, self.card, bias=bias_proj) for _ in range(dep_q)] + ) + self.to(device=device, dtype=dtype) + # We always keep the condition provider as float32. + self.condition_provider = condition_provider + self.fuser = fuser + if self.condition_provider is not None: + self.condition_provider.to(device=device) + if self.fuser is not None: + self.fuser.to(device=device) + self._init_weights() + if quantize: + replace_linear_with_qlinear(self) + + @property + def initial_token_id(self) -> int: + """Token id for the start of sequence (audio).""" + return self.card + + @property + def text_initial_token_id(self) -> int: + """Token id for the start of sequence (text).""" + return self.text_card + + @property + def text_padding_token_id(self) -> int: + """Token id for text padding.""" + return self.existing_text_padding_id + + @property + def end_of_text_padding_id(self) -> int: + """Token id for optionally marking the last padding step for a word.""" + return self.existing_text_end_padding_id + + @property + def zero_token_id(self) -> int: + """Special value in the input tokens, indicating that no sampling should + happen for that value, and no input should be given to the model.""" + return -1 + + @property + def ungenerated_token_id(self) -> int: + """Special value that can be provided in the prompt to indicate that this specific + value should be predicted and sampled. This allows for partial teacher forcing, by generating + one modality, with the other one fixed. + """ + return -2 + + @property + def device(self) -> torch.device: + first_param = next(iter(self.parameters())) + return first_param.device + + @property + def dtype(self) -> torch.dtype: + first_param = next(iter(self.text_emb.parameters())) + return first_param.dtype + + @property + def num_codebooks(self) -> int: + return self.n_q + 1 + + @property + def num_audio_codebooks(self) -> int: + return self.n_q + + @property + def audio_offset(self) -> int: + return 1 + + def _get_initial_token(self) -> torch.Tensor: + # Returns the initial token that will be fed to the model to predict the very first timestep. + # The output shape will be [B, K, 1]. + device = next(iter(self.parameters())).device + zero = torch.full( + [1, 1, 1], self.zero_token_id, device=device, dtype=torch.long + ) + special = torch.full_like(zero, self.initial_token_id) + + text_special = torch.full_like(zero, self.text_initial_token_id) + audio_token = special + text_token = text_special + audio_token = audio_token.expand(-1, self.num_audio_codebooks, -1) + token = torch.cat([text_token, audio_token], dim=1) + return token + + def forward( + self, codes: torch.Tensor, + condition_tensors: tp.Optional[ConditionTensors] = None) -> LMOutput: + """Given an input tensor of codes [B, K, T] and list of conditions, returns the logits + along with masks indicating the valid positions at which to compute the loss. + The logits time steps are aligned with those in the input `code`. + Should only be used for training, not inference (use `LMGen` for that). + + Args: + codes (torch.Tensor): Input codes of shape [B, K, T] with B the batch size, + K the number of codebooks and T the number of timesteps. When text is supported, + the first 'codebook' corresponds to the text, and the remaining codebooks are for the audio. + condition_tensors (dict[str, ConditionType], optional): pre-computed conditioning tensors. + Returns: + LMOutput: Language model outputs, containing either text or audio logits, or both. + logits (torch.Tensor, or None) of shape [B, K, T, card] corresponding to the provided codes, + i.e. the first item corresponds to logits to predict the first code, meaning that + no additional shifting of codes and logits is required. + mask (torch.Tensor, or None) of shape [B, K, T], mask over valid and invalid positions. + Given the specified interleaving strategies, parts of the logits and codes should + not be considered as valid predictions because of invalid context. + text_logits (torch.Tensor, or None) of shape [B, 1, T, text_card]. + text_mask (torch.Tensor, or None) of shape [B, 1, T], mask over the valid positions for the text. + """ + B, K, T = codes.shape + assert K == self.num_codebooks, (K, self.num_codebooks) + # Delaying codes and removing the last time step that will never be an input. + initial = self._get_initial_token().expand(B, -1, -1) + delayed_codes = _delay_sequence(self.delays, codes, initial) + # Inserting the empty tokens for the first time step. + delayed_codes = torch.cat([initial, delayed_codes], dim=2) + + sum_condition: torch.Tensor | None = None + cross_attention_src: torch.Tensor | None = None + if condition_tensors is None: + assert self.fuser is None + else: + assert self.fuser is not None + sum_condition = self.fuser.get_sum(condition_tensors) + cross_attention_src = self.fuser.get_cross(condition_tensors) + + transformer_out, text_logits = self.forward_text(delayed_codes[:, :, :-1], sum_condition, cross_attention_src) + assert transformer_out.shape[0] == delayed_codes.shape[0] + assert transformer_out.shape[1] == delayed_codes.shape[2] - 1 + logits = self.forward_depformer_training(delayed_codes[:, :, 1:], transformer_out) + + # map back the logits on pattern sequence to logits on original codes: [B, K, S, card] -> [B, K, T, card] + # and provide the corresponding mask over invalid positions of tokens. We will with NaN values invalid positions + # to ensure they properly handled. + logits, logits_mask = _undelay_sequence( + self.delays[self.audio_offset:self.audio_offset + self.dep_q], + logits, fill_value=float('NaN')) + logits_mask &= (codes[:, self.audio_offset: self.audio_offset + self.dep_q] != self.zero_token_id) + text_logits, text_logits_mask = _undelay_sequence(self.delays[:1], text_logits, fill_value=float('NaN')) + text_logits_mask &= (codes[:, :1] != self.zero_token_id) + return LMOutput(logits, logits_mask, text_logits, text_logits_mask) + + def forward_text( + self, + sequence: torch.Tensor, sum_condition: torch.Tensor | None = None, + cross_attention_src: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + B, K, S = sequence.shape + assert ( + K == self.num_codebooks + ), f"Sequence shape {sequence.shape} must match the number of codebooks." + input_sequence = sequence + input_ = None + for cb_index in range(self.num_audio_codebooks): + audio_emb = self.emb[cb_index]( + input_sequence[:, cb_index + self.audio_offset] + ) + input_ = audio_emb if input_ is None else input_ + audio_emb + text_emb = self.text_emb(input_sequence[:, 0]) + + input_ = text_emb if input_ is None else input_ + text_emb + if sum_condition is not None: + input_ = input_ + sum_condition.to(input_) + if cross_attention_src is not None: + cross_attention_src = cross_attention_src.to(input_) + transformer_out = self.transformer(input_, cross_attention_src=cross_attention_src) + if self.out_norm: + transformer_out = self.out_norm(transformer_out) + assert isinstance(transformer_out, torch.Tensor) + text_logits = self.text_linear(transformer_out) + text_logits = text_logits[:, None] + return transformer_out, text_logits + + def forward_depformer_training( + self, + sequence: torch.Tensor, + transformer_out: torch.Tensor, + ) -> torch.Tensor: + assert self.depformer_text_emb + assert self.depformer_emb + assert self.depformer + + B, K, T = sequence.shape + Ka = self.dep_q + assert ( + K == self.num_codebooks + ), f"Codebooks for Depformer training should be passed all at once, got {K}." + depformer_inputs = [] + for cb_index in range(Ka): + if self.depformer_multi_linear: + linear_index = cb_index + if self.depformer_weights_per_step_schedule is not None: + linear_index = self.depformer_weights_per_step_schedule[cb_index] + transformer_in = self.depformer_in[linear_index](transformer_out) + else: + transformer_in = self.depformer_in[0](transformer_out) + if cb_index == 0: + token_in = self.depformer_text_emb(sequence[:, 0]) + else: + token_in = self.depformer_emb[cb_index - 1](sequence[:, cb_index + self.audio_offset - 1]) + depformer_inputs.append(token_in + transformer_in) + depformer_input = torch.stack(depformer_inputs, 2) + # depformer_input is [B, T, K, depformer_dim], reshaping to [B * T, K, D] + depformer_input = depformer_input.view(B * T, Ka, -1) + depformer_output = self.depformer(depformer_input) + all_logits = [] + for cb_index in range(Ka): + logits = self.linears[cb_index](depformer_output[:, cb_index]) + all_logits.append(logits.view(B, T, -1)) + logits = torch.stack(all_logits, 1) + assert logits.dim() == 4, logits.shape # [B, Ka, T, card] + return logits + + def forward_depformer( + self, + depformer_cb_index: int, + sequence: torch.Tensor, + transformer_out: torch.Tensor, + ) -> torch.Tensor: + assert self.depformer_text_emb is not None + assert self.depformer_emb is not None + assert self.depformer is not None + B, K, S = sequence.shape + assert ( + K == 1 + ), f"Codebooks for Depformer streaming should be passed 1 by 1, got {K}." + assert ( + S == 1 + ), f"Steps for Depformer streaming should be passed 1 by 1, got {S}." + assert ( + transformer_out.shape[1] == 1 + ), "Transformer out should be a for a single step." + last_token_input: tp.Optional[torch.Tensor] = None + depformer_input = transformer_out + if self.depformer_multi_linear: + in_index = depformer_cb_index + if self.depformer_weights_per_step_schedule is not None: + in_index = self.depformer_weights_per_step_schedule[in_index] + depformer_input = self.depformer_in[in_index](depformer_input) + else: + depformer_input = self.depformer_in[0](depformer_input) + if depformer_cb_index == 0: + last_token_input = self.depformer_text_emb(sequence[:, 0]) + else: + last_token_input = self.depformer_emb[depformer_cb_index - 1]( + sequence[:, 0] + ) + assert last_token_input is not None + depformer_input = depformer_input + last_token_input + assert depformer_input.shape[1] == 1 + # depformer_input is [B, 1, depformer_dim]. + # The streaming state of the depformer ensures that the proper layer is run. + dep_output = self.depformer(depformer_input) + logits = self.linears[depformer_cb_index](dep_output) + logits = logits[:, None] + assert logits.dim() == 4, logits.shape # [B, Ka, S, card] + return logits + + def _init_weights(self): + """Initialization of the transformer module weights. + Mostly truncated gaussian, with `std = 1 / sqrt(dim_in)`. + Embeddings are also initialized with `1 / sqrt(dim)` rather than `1`. + Some layers are not going to be properly initialized: + - in_proj in MHA. + - depth transformer layers. + This is to match how our models were trained so far. + """ + + for emb_layer in self.emb: + _init_layer(emb_layer) + if self.depformer_emb is not None: + for emb_layer in self.depformer_emb: + _init_layer(emb_layer) + _init_layer(self.text_emb) + if self.depformer_text_emb is not None: + _init_layer(self.depformer_text_emb) + _init_layer(self.text_linear) + + for tr_layer in self.transformer.layers: + tr_layer.apply(_init_layer) + + for linear in self.linears: + _init_layer(linear) + + +@dataclass +class _LMGenState(State): + cache: torch.Tensor + initial: torch.Tensor + graphed_main: CUDAGraphed + graphed_depth: CUDAGraphed | None + offsets: torch.Tensor + offset_cpu: int = 0 + condition_sum: torch.Tensor | None = None + condition_cross: torch.Tensor | None = None + cfg_is_masked_until: torch.Tensor | None = None + exit_stack: ExitStack = field(default_factory=ExitStack) + reset_callback: tp.Callable[[torch.Tensor], None] | None = None + set_exec_mask_callback: tp.Callable[[torch.Tensor], None] | None = None + + def reset(self, reset_mask: torch.Tensor) -> None: + super().reset(reset_mask) + self.offsets[:] = torch.where(reset_mask, torch.zeros_like(self.offsets), self.offsets) + self.offset_cpu = 0 + if self.reset_callback is not None: + self.reset_callback(reset_mask) + + def set_exec_mask(self, exec_mask: torch.Tensor): + super().set_exec_mask(exec_mask) + if self.set_exec_mask_callback is not None: + self.set_exec_mask_callback(exec_mask) + + def __enter__(self): + self.exit_stack.__enter__() + + def __exit__(self, exc_type, exc_value, traceback): + self.exit_stack.__exit__(exc_type, exc_value, traceback) + + +class LMGen(StreamingModule[_LMGenState]): + def __init__( + self, + lm_model: LMModel, + use_sampling: bool = True, + temp: float = 0.8, + temp_text: float = 0.7, + top_k: int = 250, + top_k_text: int = 25, + cfg_coef: float = 1., + check: bool = False, + condition_tensors: ConditionTensors | None = None, + on_text_hook: tp.Optional[tp.Callable[[torch.Tensor], None]] = None, + on_text_logits_hook: tp.Optional[tp.Callable[[torch.Tensor], None]] = None, + on_audio_hook: tp.Optional[tp.Callable[[torch.Tensor], None]] = None, + support_out_of_sync: bool = False, + cfg_is_masked_until: list[int] | None = None, + cfg_is_no_text: bool = False, + ): + assert not lm_model.training, "generation shouldn't be used in training mode." + super().__init__() + + self.lm_model = lm_model + self.lm_model.set_streaming_detached(True) + self.use_sampling = use_sampling + self.temp = temp + self.temp_text = temp_text + self.top_k = top_k + self.top_k_text = top_k_text + self.cfg_coef = cfg_coef + self.check = check + self.max_delay = max( + lm_model.delays + ) # with delays, we need to generate a few more time steps. + self.delays_cuda = torch.tensor( + lm_model.delays, device=lm_model.device, dtype=torch.long + ) + self.condition_tensors = condition_tensors + self.on_text_hook = on_text_hook + self.on_text_logits_hook = on_text_logits_hook + self.on_audio_hook = on_audio_hook + self.support_out_of_sync = support_out_of_sync + self.cfg_is_masked_until = cfg_is_masked_until + self.cfg_is_no_text = cfg_is_no_text + if self.cfg_coef != 1.: + if not self.cfg_is_no_text and not self.cfg_is_masked_until: + assert self.lm_model.fuser is not None, "Model has no fuser, cannot do CFG." + assert self.condition_tensors, "Missing condition tensors for CFG." + + def _init_streaming_state(self, batch_size: int) -> _LMGenState: + lm_model = self.lm_model + initial = lm_model._get_initial_token() + cache = torch.full( + (batch_size, self.lm_model.num_codebooks, self.max_delay + 2), + lm_model.ungenerated_token_id, + device=lm_model.device, + dtype=torch.long, + ) + offsets = torch.zeros(batch_size, device=lm_model.device, dtype=torch.long) + + if self.lm_model.fuser is None: + assert not self.condition_tensors + condition_sum = None + condition_cross = None + else: + assert self.condition_tensors is not None + condition_sum = self.lm_model.fuser.get_sum(self.condition_tensors) + condition_cross = self.lm_model.fuser.get_cross(self.condition_tensors) + if condition_sum is not None: + condition_sum = condition_sum.to(self.lm_model.dtype) + if condition_cross is not None: + condition_cross = condition_cross.to(self.lm_model.dtype) + + disable = lm_model.device.type != 'cuda' + graphed_main = CUDAGraphed(lm_model.forward_text, disable=disable) + if lm_model.depformer is not None: + graphed_depth = CUDAGraphed(self.depformer_step, disable=disable) + else: + graphed_depth = None + + if self.cfg_is_masked_until is None: + cfg_is_masked_until = None + else: + cfg_is_masked_until = torch.tensor(self.cfg_is_masked_until, dtype=torch.long, device=lm_model.device) + + state = _LMGenState( + batch_size, lm_model.device, cache, initial, graphed_main, graphed_depth, + offsets, condition_sum=condition_sum, condition_cross=condition_cross, + cfg_is_masked_until=cfg_is_masked_until) + + if self.cfg_coef != 1.: + batch_size *= 2 + if state.condition_sum is not None: + assert state.condition_sum.shape[0] == batch_size, "cfg requires 2x more conditions." + if state.condition_cross is not None: + assert state.condition_cross.shape[0] == batch_size, "cfg requires 2x more conditions." + state.exit_stack.enter_context(self.lm_model.streaming(batch_size)) + + def _reset_callback(reset_mask: torch.Tensor) -> None: + if self.cfg_coef != 1.: + reset_mask = reset_mask.repeat(2) + self.lm_model.reset_streaming(reset_mask) + + def _set_exec_mask_callback(exec_mask: torch.Tensor) -> None: + if self.cfg_coef != 1.: + exec_mask = exec_mask.repeat(2) + self.lm_model.set_exec_mask(exec_mask) + + state.reset_callback = _reset_callback + state.set_exec_mask_callback = _set_exec_mask_callback + return state + + @torch.no_grad() + def _step(self, input_tokens: torch.Tensor, + depformer_replace_tokens: torch.Tensor | None = None + ) -> tuple[torch.Tensor, torch.Tensor] | None: + state = self._streaming_state + if state is None: + raise RuntimeError( + "You should wrap those calls with a `with lm_gen.streaming(): ...`." + ) + lm_model = self.lm_model + + assert input_tokens.dim() == 3, "Shape should be [B, K, T]." + B, Ki, S = input_tokens.shape + assert B == state.batch_size, f"Got a batch size {B}, expected {state.batch_size}" + assert S == 1, "Only support being given steps one by one." + needed_tokens = lm_model.num_codebooks - lm_model.dep_q - 1 + assert ( + Ki >= needed_tokens + ), f"We expect {needed_tokens} tokens from the user stream, got {Ki}." + + if Ki > needed_tokens: + input_tokens = input_tokens[:, :needed_tokens, :] + + CT = state.cache.shape[2] + + delays = self.delays_cuda[lm_model.dep_q + 1:] + write_positions = (state.offsets[:, None, None] + delays[:, None]) % CT + scatter_with_mask_(state.cache[:, lm_model.dep_q + 1:], -1, write_positions, input_tokens, + state.exec_mask[:, None, None]) + + is_init = state.offsets[:, None, None] <= self.delays_cuda[:, None] + is_init |= ~state.exec_mask[:, None, None] # we also give init tokens if not executing to avoid crashing. + positions = (state.offsets % CT)[:, None, None].expand_as(is_init) + input_ = state.cache.gather(dim=2, index=positions) + input_ = torch.where(is_init, state.initial, input_) + + if self.check: + # Check that we are not feeding in any value that is not generated yet. + assert not (input_ == lm_model.ungenerated_token_id).any(), ( + state.offsets, + input_, + ) + assert (input_[:, lm_model.audio_offset :] <= lm_model.card).all(), input_ + assert (input_[:, :1] <= lm_model.text_card).all() + + zero = torch.full((1,), self.lm_model.zero_token_id, dtype=torch.long, device=input_.device) + if self.cfg_coef != 1.: + if state.cfg_is_masked_until is not None: + limit = self.delays_cuda[:, None] + state.cfg_is_masked_until.view(-1, 1, 1) + is_zeroed = state.offsets[:, None, None] <= limit + + masked = torch.where(is_zeroed & ~is_init, zero, input_) + input_ = torch.cat([input_, masked], dim=0) + else: + input_ = input_.repeat(2, 1, 1) + if self.cfg_is_no_text: + input_[B:, :1] = torch.where(~is_init[:, :1], zero, input_[B:, :1]) + + transformer_out, text_logits = state.graphed_main(input_, state.condition_sum, state.condition_cross) + if self.cfg_coef != 1.: + logits, logits_null = text_logits.chunk(2) + if self.cfg_is_no_text: + text_logits = logits + else: + text_logits = logits_null + (logits - logits_null) * self.cfg_coef + # Shape of text_logits should be [B, K_text=1, T=1, Card_text] + if self.on_text_logits_hook: + self.on_text_logits_hook(text_logits) + text_token = sample_token( + text_logits.float(), + self.use_sampling, + self.temp_text, + self.top_k_text, + ) + assert text_token.dim() == 3, text_token.shape + assert text_token.shape[2] == 1 + assert text_token.shape[1] == 1, "Only one text stream supported." + text_token = text_token[:, 0, 0] # shape is [B] + if self.on_text_hook is not None: + self.on_text_hook(text_token) + if state.graphed_depth is None: + audio_tokens = None + elif depformer_replace_tokens is None: + audio_tokens = state.graphed_depth(text_token, transformer_out) + if self.on_audio_hook is not None: + self.on_audio_hook(audio_tokens) + else: + assert depformer_replace_tokens.dim() == 3 + audio_tokens = depformer_replace_tokens.squeeze(-1) + + state.offsets = torch.where(state.exec_mask, state.offsets + 1, state.offsets) + state.offset_cpu += 1 + positions = (state.offsets % CT)[:, None, None] + scatter_with_mask_(state.cache[:, :1], -1, positions, + text_token[:, None, None], state.exec_mask[:, None, None]) + if audio_tokens is not None: + audio_tokens = audio_tokens[:, :, None] + scatter_with_mask_( + state.cache[:, 1 : lm_model.dep_q + 1, :], + -1, + positions.expand_as(audio_tokens), + audio_tokens, + state.exec_mask[:, None, None], + ) + + if not self.support_out_of_sync and state.offset_cpu <= self.max_delay: + # When using out of sync exec, should not rely on this being None. + return None + B = state.cache.shape[0] + gen_delays_cuda = self.delays_cuda[: lm_model.dep_q + 1] + index = (state.offsets[:, None, None] - self.max_delay + gen_delays_cuda[:, None]) % CT + out = state.cache.gather(dim=2, index=index) + mask = (state.offsets <= self.max_delay) | ~state.exec_mask + out[mask, :, :] = lm_model.ungenerated_token_id + return out, transformer_out + + @torch.no_grad() + def step(self, input_tokens: torch.Tensor, + depformer_replace_tokens: torch.Tensor | None = None) -> torch.Tensor | None: + out = self._step(input_tokens, depformer_replace_tokens) + if out is None: + return None + return out[0] + + @torch.no_grad() + def step_with_extra_heads( + self, + input_tokens: torch.Tensor, + depformer_replace_tokens: torch.Tensor | None = None, + ) -> tuple[torch.Tensor, list[torch.Tensor]] | None: + out = self._step(input_tokens, depformer_replace_tokens) + if out is None: + return None + out, transformer_out = out + extra_heads = [extra_head(transformer_out) for extra_head in self.lm_model.extra_heads] + return out, extra_heads + + def depformer_step( + self, + text_token: torch.Tensor, + transformer_out: torch.Tensor, + ) -> torch.Tensor: + B, = text_token.shape + B_cfg = B + if self.cfg_coef != 1.: + B_cfg = 2 * B + prev_token = text_token + lm_model = self.lm_model + depformer_tokens: list[torch.Tensor] = [] + assert lm_model.depformer + assert not lm_model.depformer.is_streaming + with lm_model.depformer.streaming(B_cfg): + assert lm_model.depformer.is_streaming + for cb_index in range(lm_model.dep_q): + input_ = prev_token[:, None, None] + if self.cfg_coef != 1.: + input_ = input_.repeat(2, 1, 1) + logits = lm_model.forward_depformer(cb_index, input_, transformer_out) + if self.cfg_coef != 1.: + logits, logits_null = logits.chunk(2) + logits = logits_null + (logits - logits_null) * self.cfg_coef + next_token = sample_token( + logits.float(), + self.use_sampling, + self.temp, + self.top_k, + ) + assert next_token.shape == (B, 1, 1) + next_token = next_token[:, 0, 0] # shape is B + depformer_tokens.append(next_token) + prev_token = next_token + + assert len(depformer_tokens) == lm_model.dep_q, ( + len(depformer_tokens), + lm_model.dep_q, + ) + out = torch.stack(depformer_tokens, dim=1) + assert out.shape == (B, lm_model.dep_q), out.shape + return out diff --git a/moshi_src/moshi/models/lm_utils.py b/moshi_src/moshi/models/lm_utils.py new file mode 100644 index 0000000..de96ebf --- /dev/null +++ b/moshi_src/moshi/models/lm_utils.py @@ -0,0 +1,124 @@ +import math +import typing as tp +import torch +from torch import nn + +from ..modules.transformer import create_norm_fn + + +def _delay_sequence(delays: tp.List[int], tensor: torch.Tensor, padding: torch.Tensor) -> torch.Tensor: + B, K, T = tensor.shape + assert len(delays) == K, (len(delays), K) + outs = [] + + for k, delay in enumerate(delays): + assert delay >= 0 + line = tensor[:, k].roll(delay, dims=1) + if delay > 0: + line[:, :delay] = padding[:, k] + outs.append(line) + return torch.stack(outs, dim=1) + + +def _undelay_sequence(delays: tp.List[int], tensor: torch.Tensor, + fill_value: tp.Union[int, float] = float('NaN')) -> tp.Tuple[torch.Tensor, torch.Tensor]: + B, K, T, *_ = tensor.shape + assert len(delays) == K + mask = torch.ones(B, K, T, dtype=torch.bool, device=tensor.device) + outs = [] + if all([delay == 0 for delay in delays]): + return tensor, mask + for k, delay in enumerate(delays): + assert delay >= 0 + line = tensor[:, k].roll(-delay, dims=1) + if delay > 0: + line[:, -delay:] = fill_value + mask[:, k, -delay:] = 0 + outs.append(line) + return torch.stack(outs, dim=1), mask + + +def _get_init_fn(input_dim: int) -> tp.Callable[[torch.Tensor], None]: + def _init(x: torch.Tensor) -> None: + std = 1 / math.sqrt(input_dim) + x_orig = x + if x.device.type == 'cpu' and x.dtype in [torch.float16, torch.bfloat16]: + x = x.float() + + torch.nn.init.trunc_normal_(x, mean=0.0, std=std, a=-3 * std, b=3 * std) + if x_orig is not x: + x_orig.data[:] = x.to(x_orig) + return _init + + +def _init_layer(m: nn.Module, + zero_bias_init: bool = True): + if isinstance(m, nn.Linear): + init_fn = _get_init_fn(m.in_features) + init_fn(m.weight) + if zero_bias_init and m.bias is not None: + nn.init.constant_(m.bias, 0) + elif isinstance(m, nn.Embedding): + init_fn = _get_init_fn(m.embedding_dim) + init_fn(m.weight) + + +class ScaledEmbedding(nn.Embedding): + """Boost learning rate for embeddings (with `scale`). + + Args: + norm (bool): if True, uses a layer norm after the embedding. + zero_idx (int): special value indicating that the output should be exactly 0. + low_rank (int | None): if provided, uses low rank embedding with a linear layer to reach + the desired dimension. Quite efficient for reducing the number of weights for very large vocabs. + lr (float or None): learning rate to use, only valid if the `make_optim_group()` method is used. + demux_second_stream (bool): input tokens can be the cartesian product of the vocab size, + and they will be demuxed, e.g. `(tok2 * card + tok1)`. In that case the same embedding + is used with different linear matrices. + """ + + def __init__(self, num_embeddings: int, embedding_dim: int, + *args, norm: bool = False, zero_idx: int = -1, + low_rank: int | None = None, lr: float | None = None, + demux_second_stream: bool = False, **kwargs): + super().__init__(num_embeddings, low_rank or embedding_dim, *args, **kwargs) + self.norm = None + if norm: + self.norm = create_norm_fn("layer_norm", self.embedding_dim) + assert zero_idx < 0, "Please use negative values for the zero_idx." + self.zero_idx = zero_idx + self.lr = lr + self.low_rank = None + if low_rank is not None: + self.low_rank = nn.Linear(low_rank, embedding_dim, bias=False) + + self.demux_second_stream = demux_second_stream + assert self.zero_idx == -1, "When demuxing a second stream, zero_idx must be -1." + if self.demux_second_stream: + assert not norm + self.out1 = nn.Linear(low_rank or embedding_dim, embedding_dim, bias=False) + self.out2 = nn.Linear(low_rank or embedding_dim, embedding_dim, bias=False) + + def forward(self, input, *args, **kwargs): + is_zero = input == self.zero_idx + zero = torch.zeros(1, dtype=input.dtype, device=input.device) + input = input.clamp(min=0) + if self.demux_second_stream: + left = input % self.num_embeddings + right = input // self.num_embeddings + # Right is itself between [-1, ..., card - 1], with -1 being the zero value. + right = right - 1 + left = super().forward(left, *args, **kwargs) + right_zero = (right < 0)[..., None] + right.clamp_(min=0) + right = super().forward(right, *args, **kwargs) + y = self.out1(left) + torch.where(right_zero, zero, self.out2(right)) + y = torch.where(is_zero[..., None], zero, y) + else: + y = super().forward(input, *args, **kwargs) + if self.norm is not None: + y = self.norm(y) + y = torch.where(is_zero[..., None], zero, y) + if self.low_rank is not None: + y = self.low_rank(y) + return y diff --git a/moshi_src/moshi/models/loaders.py b/moshi_src/moshi/models/loaders.py new file mode 100644 index 0000000..7aa7d31 --- /dev/null +++ b/moshi_src/moshi/models/loaders.py @@ -0,0 +1,481 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +"""Retrieves the pretrained models for Moshi and Mimi.""" + +from dataclasses import dataclass, field +import json +from pathlib import Path +import warnings +from huggingface_hub import hf_hub_download + +try: + from huggingface_hub.errors import EntryNotFoundError +except ImportError: + from huggingface_hub.utils import EntryNotFoundError # pyright: ignore +from safetensors.torch import load_model, load_file +import sentencepiece +import torch +import typing as tp +from .compression import MimiModel +from ..conditioners import BaseConditioner, ConditionProvider, ConditionFuser +from .lm import LMModel +from ..modules import SEANetEncoder, SEANetDecoder, transformer +from ..quantization import SplitResidualVectorQuantizer +from ..modules.lora import replace_all_linear_with_lora, replace_lora_with_linear + + +SAMPLE_RATE = 24000 +FRAME_RATE = 12.5 + +TEXT_TOKENIZER_NAME = "tokenizer_spm_32k_3.model" +MOSHI_NAME = "model.safetensors" +MOSHI_Q8_NAME = "model.q8.safetensors" +MIMI_NAME = "tokenizer-e351c8d8-checkpoint125.safetensors" +DEFAULT_REPO = "kyutai/moshiko-pytorch-bf16" + + +_seanet_kwargs = { + "channels": 1, + "dimension": 512, + "causal": True, + "n_filters": 64, + "n_residual_layers": 1, + "activation": "ELU", + "compress": 2, + "dilation_base": 2, + "disable_norm_outer_blocks": 0, + "kernel_size": 7, + "residual_kernel_size": 3, + "last_kernel_size": 3, + # We train using weight_norm but then the weights are pre-processed for inference so + # that we can use a normal convolution. + "norm": "none", + "pad_mode": "constant", + "ratios": [8, 6, 5, 4], + "true_skip": True, +} +_quantizer_kwargs = { + "dimension": 256, + "n_q": 32, + "bins": 2048, + "input_dimension": _seanet_kwargs["dimension"], + "output_dimension": _seanet_kwargs["dimension"], +} +_transformer_kwargs = { + "d_model": _seanet_kwargs["dimension"], + "num_heads": 8, + "num_layers": 8, + "causal": True, + "layer_scale": 0.01, + "context": 250, + "conv_layout": True, + "max_period": 10000, + "gating": "none", + "norm": "layer_norm", + "positional_embedding": "rope", + "dim_feedforward": 2048, + "input_dimension": _seanet_kwargs["dimension"], + "output_dimensions": [_seanet_kwargs["dimension"]], +} + +_lm_kwargs = { + "dim": 4096, + "text_card": 32000, + "existing_text_padding_id": 3, + "n_q": 16, + "dep_q": 8, + "card": _quantizer_kwargs["bins"], + "num_heads": 32, + "num_layers": 32, + "hidden_scale": 4.125, + "causal": True, + "layer_scale": None, + "context": 3000, + "max_period": 10000, + "gating": "silu", + "norm": "rms_norm_f32", + "positional_embedding": "rope", + "depformer_dim": 1024, + "depformer_dim_feedforward": int(4.125 * 1024), + "depformer_num_heads": 16, + "depformer_num_layers": 6, + "depformer_layer_scale": None, + "depformer_multi_linear": True, + "depformer_context": 8, + "depformer_max_period": 10000, + "depformer_gating": "silu", + "depformer_pos_emb": "none", + "depformer_weights_per_step": True, + "delays": [0, 0, 1, 1, 1, 1, 1, 1, 1, 0, 1, 1, 1, 1, 1, 1, 1], +} + + +def hf_get(filename: str | Path, hf_repo: str | None = None, + check_local_file_exists: bool = False) -> Path: + if isinstance(filename, Path): + return filename + if filename.startswith("hf://"): + parts = filename.removeprefix("hf://").split("/") + repo_name = parts[0] + "/" + parts[1] + filename = "/".join(parts[2:]) + return Path(hf_hub_download(repo_name, filename)) + elif filename.startswith("file://"): + # Provide a way to force the read of a local file. + filename = filename.removeprefix("file://") + return Path(filename) + elif hf_repo is not None: + if check_local_file_exists: + if Path(filename).exists(): + return Path(filename) + return Path(hf_hub_download(hf_repo, filename)) + else: + return Path(filename) + + +@dataclass +class CheckpointInfo: + """ + Contains the paths to each sub model, along with some extra configuration. + + Args: + moshi_weights: path to the checkpoint for the Moshi LM. + mimi_weights: path to the checkpoint for the Mimi audio tokenizer. + tokenizer: path to the text tokenizer. + lm_config: config for instantiating the LM model. + Can be None if the original Moshi 7B config should be used. + raw_config: raw config, including original keys not intended for the LM. + model_type: indicate the intended use, should be `moshi` or `hibiki`. + lora_weights: path to an optional checkpoint with lora weights. + lm_gen_config: optional default params to use for generation with this model. + tts_config: optional TTS specific configuration. + stt_config: optional STT specific configuration. + model_id: optional dict containing tracability information on the model origin, in particular + its signature and epoch. + """ + + moshi_weights: Path + mimi_weights: Path + tokenizer: Path + lm_config: dict | None = None + raw_config: dict | None = None + model_type: str = "moshi" + lora_weights: Path | None = None + lm_gen_config: dict = field(default_factory=dict) + tts_config: dict = field(default_factory=dict) + stt_config: dict = field(default_factory=dict) + model_id: dict = field(default_factory=dict) + + @staticmethod + def from_hf_repo( + hf_repo: str, + moshi_weights: Path | str | None = None, + mimi_weights: Path | str | None = None, + tokenizer: Path | str | None = None, + config_path: Path | str | None = None, + lora_weights: Path | str | None = None, + ) -> "CheckpointInfo": + """Downloads the checkpoints from the given repo, along with its config. + + Extra overrides are possible for each of Moshi, Mimi, or the text tokenizer, + which should be either a Path to a local file or a string representing a path + to a local file or starting with `hf://` for pointing to a file in another repo. + + Finally, a `config_path` can be provided to override the config from the repository. + """ + if config_path is None: + try: + config_path = hf_hub_download(hf_repo, "config.json") + except EntryNotFoundError: + # No config.json, which might indicate legacy repository. + warnings.warn( + f"Repository {hf_repo} contains no config.json. " + "Assuming this is a Moshi 7B. Support for such repository " + "might be removed in the future." + ) + if config_path is None: + moshi_name = MOSHI_NAME + mimi_name = MIMI_NAME + tokenizer_name = TEXT_TOKENIZER_NAME + lm_config = None + raw_config = None + model_type = "moshi" + lm_gen_config = {} + tts_config = {} + stt_config = {} + model_id = {} + lora_name = None + else: + raw_config = json.loads(Path(config_path).read_text()) + lm_config = dict(raw_config) + moshi_name = lm_config.pop("moshi_name", MOSHI_NAME) + mimi_name = lm_config.pop("mimi_name", MIMI_NAME) + tokenizer_name = lm_config.pop("tokenizer_name", TEXT_TOKENIZER_NAME) + lora_name = lm_config.pop("lora_name", None) + model_type = lm_config.pop("model_type", "moshi") + lm_gen_config = lm_config.pop("lm_gen_config", {}) + tts_config = lm_config.pop("tts_config", {}) + stt_config = lm_config.pop("stt_config", {}) + model_id = lm_config.pop("model_id", {}) + + if moshi_weights is None: + moshi_weights_final = hf_get(moshi_name, hf_repo) + else: + moshi_weights_final = hf_get(moshi_weights) + + if mimi_weights is None: + mimi_weights_final = hf_get(mimi_name, hf_repo) + else: + mimi_weights_final = hf_get(mimi_weights) + + if tokenizer is None: + tokenizer_final = hf_get(tokenizer_name, hf_repo) + else: + tokenizer_final = hf_get(tokenizer) + + if lora_weights is None and lora_name: + lora_weights_final = hf_get(lora_name, hf_repo) + elif lora_weights is not None: + lora_weights_final = hf_get(lora_weights) + else: + lora_weights_final = None + + return CheckpointInfo( + moshi_weights_final, + mimi_weights_final, + tokenizer_final, + lm_config, + raw_config, + model_type, + lora_weights_final, + lm_gen_config=lm_gen_config, + tts_config=tts_config, + stt_config=stt_config, + model_id=model_id, + ) + + def get_mimi(self, device: torch.device | str = "cpu") -> MimiModel: + if self.lm_config is None: + num_codebooks = 8 + else: + num_codebooks = max(self.lm_config["dep_q"], self.lm_config["n_q"] - self.lm_config["dep_q"]) + return get_mimi(self.mimi_weights, num_codebooks=num_codebooks, device=device) + + def get_moshi( + self, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.bfloat16, + load_weight: bool = True, + **kwargs, + ) -> LMModel: + model = get_moshi_lm( + self.moshi_weights if load_weight else None, + lm_kwargs=self.lm_config, + device=device, + dtype=dtype, + lora_weights=self.lora_weights, + **kwargs, + ) + if self.model_type == "hibiki": + # Sometime the model samples the EOS (2) too early, which we want to ignore. + # We keep generating if the input file is not finished, and this is a way + # to implicitely replace early EOS with PAD. + model.text_emb.weight.data[2] = model.text_emb.weight.data[3] + return model + + def get_text_tokenizer(self) -> sentencepiece.SentencePieceProcessor: + return sentencepiece.SentencePieceProcessor(str(self.tokenizer)) # type: ignore + + +def _is_safetensors(path: Path | str) -> bool: + return Path(path).suffix in (".safetensors", ".sft", ".sfts") + + +def get_mimi( + filename: str | Path | None, device: torch.device | str = "cpu", num_codebooks: int = 8 +) -> MimiModel: + """Return a pretrained Mimi model, or unintialized if `filename` is None.""" + encoder = SEANetEncoder(**_seanet_kwargs) + decoder = SEANetDecoder(**_seanet_kwargs) + encoder_transformer = transformer.ProjectedTransformer( + device=device, **_transformer_kwargs + ) + decoder_transformer = transformer.ProjectedTransformer( + device=device, **_transformer_kwargs + ) + quantizer = SplitResidualVectorQuantizer( + **_quantizer_kwargs, + ) + model = MimiModel( + encoder, + decoder, + quantizer, + channels=1, + sample_rate=SAMPLE_RATE, + frame_rate=FRAME_RATE, + encoder_frame_rate=SAMPLE_RATE / encoder.hop_length, + causal=True, + resample_method="conv", + encoder_transformer=encoder_transformer, + decoder_transformer=decoder_transformer, + ).to(device=device) + model.eval() + if filename is not None: + if _is_safetensors(filename): + load_model(model, filename, device=str(device)) + else: + pkg = torch.load(filename, "cpu") + model.load_state_dict(pkg["model"]) + model.set_num_codebooks(num_codebooks) + return model + + +def get_moshi_lm( + filename: str | Path | None, + lm_kwargs: tp.Optional[tp.Dict[str, tp.Any]] = None, + device: torch.device | str = "cpu", + dtype: torch.dtype = torch.bfloat16, + lora_weights: str | Path | None = None, + fuse_lora: bool = False, + lm_kwargs_overrides={}, +) -> LMModel: + if lm_kwargs is None: + lm_kwargs = _lm_kwargs + lm_kwargs = dict(lm_kwargs) + assert lm_kwargs is not None + + if "conditioners" in lm_kwargs: + lm_kwargs["condition_provider"] = get_conditioner_provider( + lm_kwargs["dim"], device, lm_kwargs + ) + del lm_kwargs["conditioners"] + if "fuser" in lm_kwargs: + lm_kwargs["fuser"] = get_condition_fuser(lm_kwargs) + + lm_kwargs = lm_kwargs | lm_kwargs_overrides + assert lm_kwargs is not None + + # deprecated params. + lm_kwargs.pop("depformer_causal", None) + + # moved params + if 'demux_second_stream' in lm_kwargs: + lm_kwargs['demux_second_text_stream'] = lm_kwargs.pop('demux_second_stream') + + # lora params. + lora = lm_kwargs.pop("lora", False) + lora_rank = lm_kwargs.pop("lora_rank", 128) + lora_scaling = lm_kwargs.pop("lora_scaling", 2.0) + + init_device = device + if filename is not None: + init_device = torch.device('meta') + + model = LMModel( + device=init_device, + dtype=dtype, + **lm_kwargs) + + if filename is not None: + if _is_safetensors(filename): + state = load_file(filename, device=str(device)) + for key, value in state.items(): + if value.dtype.is_floating_point: + if key.startswith('condition_provider.') or key.startswith('fuser.'): + value = value.float() + else: + value = value.to(dtype) + state[key] = value + model.load_state_dict(state, assign=True) + + else: + pkg = torch.load(filename, "cpu",) + model.load_state_dict(pkg["fsdp_best_state"]["model"], assign=True) + + if lora: + assert not lm_kwargs.get("quantize"), ( + "LoRA and quantization are incompatible for now." + ) + model = get_lora_moshi( + model=model, + lora_rank=lora_rank, + lora_scaling=lora_scaling, + lora_weights=lora_weights, + device=device, + dtype=dtype, + fuse_lora=fuse_lora, + ) + else: + assert lora_weights is None, ( + "`lora` is False, but received some lora_weights to load." + ) + model.eval() + return model + + +def get_conditioner( + output_dim: int, device: torch.device | str, conditioner_cfg: dict +) -> BaseConditioner: + conditioner_type = conditioner_cfg["type"] + conditioner_kwargs = conditioner_cfg[conditioner_type] + conditioner_kwargs.update({"output_dim": output_dim, "device": device}) + if conditioner_type == "lut": + from ..conditioners.text import LUTConditioner + return LUTConditioner(**conditioner_kwargs) + elif conditioner_type == "tensor": + from ..conditioners.tensors import TensorConditioner + return TensorConditioner(**conditioner_kwargs) + else: + raise RuntimeError(f"Unknow conditioner type {conditioner_type}.") + + +def get_conditioner_provider( + output_dim: int, device: torch.device | str, cfg: dict +) -> ConditionProvider: + """Instantiate a conditioning model.""" + conditioners: tp.Dict[str, BaseConditioner] = {} + for cond, cond_cfg in cfg["conditioners"].items(): + conditioners[cond] = get_conditioner(output_dim, device, cond_cfg) + conditioner = ConditionProvider(conditioners, device=device) + return conditioner + + +def get_condition_fuser(cfg: dict) -> ConditionFuser: + """Instantiate a condition fuser object.""" + fuser_cfg = cfg["fuser"] + fuser_methods = ["sum", "cross", "prepend"] + fuse2cond = {k: fuser_cfg.get(k, []) for k in fuser_methods} + kwargs = {k: v for k, v in fuser_cfg.items() if k not in fuser_methods} + fuser = ConditionFuser(fuse2cond=fuse2cond, **kwargs) + return fuser + + +def get_lora_moshi( + model: LMModel, + lora_weights: str | Path | None, + lora_rank: int, + lora_scaling: float, + dtype: torch.dtype = torch.bfloat16, + device: torch.device | str = "cpu", + fuse_lora: bool = True, +) -> LMModel: + init_device = device + if lora_weights is not None: + init_device = torch.device('meta') + replace_all_linear_with_lora(model, lora_rank, lora_scaling, device=init_device) + if lora_weights is not None: + assert _is_safetensors(lora_weights), "LoRA weights must be a safetensors file." + lora_state_dict = load_file(lora_weights, device=str(device)) + for key, value in lora_state_dict.items(): + if value.dtype.is_floating_point: + value = value.to(dtype=dtype) + lora_state_dict[key] = value + res = model.load_state_dict(lora_state_dict, strict=False, assign=True) + if res.unexpected_keys: + raise RuntimeError( + f"unexpected_keys in the lora weights: {res.unexpected_keys}" + ) + model = model.to(dtype=dtype, device=device) + if fuse_lora: + replace_lora_with_linear(model) + return model diff --git a/moshi_src/moshi/models/tts.py b/moshi_src/moshi/models/tts.py new file mode 100644 index 0000000..124291e --- /dev/null +++ b/moshi_src/moshi/models/tts.py @@ -0,0 +1,621 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +""" +Implements the logic for the state machine around the Delayed Streams Modeling (DSM) based TTS model. +Things are more complex than for STT models, where we can simply force feed the audio tokens and +sample the text ones. For TTS we start from pure text, not text properly padded with time alignment. +We co-generate the padded text sequence along with the audio output from the original text by +having the model signal us when it thinks the next step will be the start of a word. We then pop a word +to feed and feed it the token representation of the word over the next few steps. +""" + +from collections import deque +from dataclasses import dataclass, field +from functools import cached_property +import re +from pathlib import Path +import typing as tp + +from safetensors.torch import load_file +from sentencepiece import SentencePieceProcessor +import sphn +import torch + +from ..conditioners import ConditionAttributes, dropout_all_conditions, TensorCondition +from ..conditioners.text import LUTConditioner +from . import loaders, MimiModel, LMModel, LMGen + + +DEFAULT_DSM_TTS_REPO = 'kyutai/tts-1.6b-en_fr' +DEFAULT_DSM_TTS_VOICE_REPO = 'kyutai/tts-voices' + + +@dataclass +class TokenIds: + """ + The token ids for special tokens: + - card: text cardinality, including the initial token (1 + tokenizer cardinality). + This is used for multiplexing multiple input tokens into the text stream. + - new_word: a new word is starting. + - pad: padding, nothing happens. + - main: indicates the start of turn of the main speaker. + - other: indicates the start of turn of the other speaker. + - zero: special value that is embedded to exactly 0. + - ungenerated: indicate that a value is not yet generated but should be + + """ + card: int + new_word: int = 0 + pad: int = 3 + main: int = 1 + other: int = 2 + zero: int = -1 + ungenerated: int = -2 + + +@dataclass +class Entry: + """One word to generate. + + Args: + tokens: list of tokens for this word. + text: word as string. + padding: if > 0, we will prevent the model from sampling a new word for that + many steps after the current word. Note that even for `padding=0`, the model + will be forbidden to sample a new word until all the tokens for the current word are consumed. + audio_tokens: is used when some audio should be used as a prefix in the model.""" + tokens: list[int] + text: str + padding: int = 0 + audio_tokens: torch.Tensor | None = None + + +@dataclass +class State: + """State of the TTS Machine. + + Args: + entries: queue containing the entries to generate. + remaining_padding: how many times the model can still sample a pad. + forced_padding: how many times the model is still forced to sample a pad. + queued: queue containing the main stream text tokens to feed. + lookahead_queued: queue containing the lookahead text tokens to feed. + end_step: once we reach the end of the generation, this is set to the current step. + The end of the generation is once the model samples a `word` but `entries` is empty. + consumption_times: list of steps at which each entry in `entries` was consumed. + transcript: list of tuples `(word, step)`, at which each word was consumed. + audio_tokens_remaining: when using some audio as prefix, this would contain the remaining + audio tokens to force into the model. + zero_text_remaining: when using an audio only prefix, for how long we should still force + the text to be `zero`. + """ + entries: deque[Entry] + remaining_padding: int + forced_padding: int + queued: deque[int] = field(default_factory=deque) + lookahead_queued: deque[int] = field(default_factory=deque) + end_step: int | None = None + consumption_times: list[int] = field(default_factory=list) + transcript: list[tuple[str, int]] = field(default_factory=list) + + def get_tokens_ahead(self, lookahead: int) -> list[int]: + assert lookahead > 0 + for entry in self.entries: + if entry.tokens: + lookahead -= 1 + if lookahead == 0: + return entry.tokens + return [] + + +def _delayed(codes: torch.Tensor, delays: list[int], fill_value: int) -> torch.Tensor: + # Apply the acoustic delay on the provided audio tokens. + K, T = codes.shape + out = torch.full((K, T + max(delays)), fill_value, device=codes.device, dtype=torch.long) + for k, delay in enumerate(delays): + out[k, delay: delay + T] = codes[k] + return out + + +def _make_null(all_attributes: tp.Sequence[ConditionAttributes]) -> list[ConditionAttributes]: + # When using CFG, returns the null conditions. + return dropout_all_conditions(all_attributes) + + +@dataclass +class StateMachine: + """State machine that manipulates the `State` based on the model prediction. + In particular, every time the model predicts a `word` (see `TokenIds`) special token, + the state machine will pop the next word to synthesize and start feeding it. + The model is optionally equipped with a second input text stream providing a lookahead + into the future text. + + Args: + token_ids: special token values. + second_stream_ahead: if > 0, the model needs a second stream for lookahead. + max_padding: maximum number of padding tokens that can be sampled in a row. + initial_padding: number of padding tokens at the beginning, to prevent the first + word from being cut. + + """ + + token_ids: TokenIds + second_stream_ahead: int = 0 + max_padding: int = 6 + initial_padding: int = 2 + + def new_state(self, entries: tp.Sequence[Entry]) -> State: + state = State( + entries=deque(entries), + lookahead_queued=deque(), + remaining_padding=self.initial_padding, + forced_padding=self.initial_padding, + ) + return state + + def process(self, step: int, state: State, token: int) -> tuple[int, bool]: + """ + Process the output of the model. + + Args: + step: current step index. + state: state to act upon. + token: model prediction + + Returns: + - output_token: value to use as the text input for the model at the next step. + - consumed_new_word: True if a new word was consumed. + """ + consumed_new_word = False + if token not in [self.token_ids.new_word, self.token_ids.pad]: + token = self.token_ids.pad + + if state.queued: + # Some text tokens are yet to be fed, we must PAD. + token = self.token_ids.pad + elif state.forced_padding > 0: + # We are forced to pad, we must PAD. + token = self.token_ids.pad + elif state.remaining_padding <= 0: + # We are not allowed to pad, we must ask for a new WORD. + token = self.token_ids.new_word + + if token == self.token_ids.new_word: + if state.entries: + entry = state.entries.popleft() + state.consumption_times.append(step) + if entry.tokens: + consumed_new_word = True + state.transcript.append((entry.text, step)) + # We queue the tokens to be fed to the model. + state.queued.extend(entry.tokens) + if self.second_stream_ahead: + # We queue the tokens for the N+lookahead word into the second text stream. + state.lookahead_queued.extend(state.get_tokens_ahead(self.second_stream_ahead)) + # Entry contains a new word, we reset the max padding counter. + state.remaining_padding = self.max_padding + else: + # Entry is only here to insert a break, pretend the token was a PAD. + token = self.token_ids.pad + state.forced_padding = entry.padding + else: + token = self.token_ids.pad + if self.second_stream_ahead and state.end_step is None: + token = self.token_ids.new_word + # Trying to consume past the last word, we reached the end. + if state.end_step is None: + state.end_step = step + + output: int | None = None + if token == self.token_ids.pad: + # Decrement the counters for remaining and forced pads. + if state.remaining_padding > 0: + state.remaining_padding -= 1 + if state.forced_padding > 0: + state.forced_padding -= 1 + if state.queued: + # We have some text tokens to feed to the model. + output = state.queued.popleft() + else: + output = self.token_ids.pad + elif token == self.token_ids.new_word: + output = self.token_ids.new_word + elif token == self.token_ids.zero: + output = token + else: + raise RuntimeError(f"Invalid token {token}") + + if self.second_stream_ahead: + second = -1 + if output == self.token_ids.new_word: + # If sampled the `word` special token, we put it on the + # second text stream instead of the main one. + second = self.token_ids.new_word + if state.queued: + # This allows us to pass the current word tokens faster. + output = state.queued.popleft() + else: + output = self.token_ids.pad + elif state.lookahead_queued: + # Otherwise if we have some lookahead tokens we feed them. + second = state.lookahead_queued.popleft() + # Then we multiplex the two tokens. We add `+1` to `second` so that + # we can encode -1, which would translate to an all 0s embedding. + # This will get de-multiplexed in the embedding in lm.py. + output = (second + 1) * self.token_ids.card + output + + assert output is not None + return output, consumed_new_word + + +def script_to_entries(tokenizer: SentencePieceProcessor, token_ids: TokenIds, frame_rate: float, + script: tp.Sequence[str], multi_speaker: bool = True, padding_between: int = 0) -> list[Entry]: + """Process a given script into a list of `Entry` that will be consumed by the model. + + This function will perform some replacements such as removing some caracters such as ':', etc. + It also supports a single XML tag from SSML, namely ``. This allow the insertion + of a pause of roughly the requested duration. + + Args: + tokenizer: text tokenizer. + token_ids: See `TokenIds`. + frame_rate: frame rate of the audio codec. + script: list of turns, with each element indicating a change of turn. Starts with the main speaker. + Use an empty first turn to start with the other speaker. + multi_speaker: whether the model was trained to handle more than one speaker. + padding_between: amount of padding to force between words. Will make the model articulate + a bit better with values such as 1. + """ + speaker_tokens = [token_ids.main, token_ids.other] + last_speaker = None + entries = [] + + # break is indicated as e.g. + event_re = re.compile(r"(?:)|(?:\s+)") + + def _add_entry(idx: int, word: str): + nonlocal first_content, last_speaker + assert ' ' not in word + assert word + tokens = tokenizer.encode(word) # type: ignore + if first_content: + speaker = idx % len(speaker_tokens) + if multi_speaker and last_speaker != speaker: + last_speaker = speaker + tokens.insert(0, speaker_tokens[speaker]) + first_content = False + padding = 0 + if padding_between > 0: + padding = max(0, padding_between + len(tokens) - 1) + entries.append(Entry(tokens=tokens, text=word, padding=padding)) + + for idx, line in enumerate(script): + first_content = True + line = line.replace('’', "'") + line = line.replace(':', " ") + line = line.replace('(', "") + line = line.replace(')', "") + while line: + match = event_re.search(line) + if match is None: + break + word = line[:match.start()] + line = line[match.end():] + if word: + _add_entry(idx, word) + if match.group(1): + break_duration = float(match.group(1)) + padding = int(round(break_duration * frame_rate)) + entry = Entry(tokens=[], text='', padding=padding) + entries.append(entry) + if line: + _add_entry(idx, line) + return entries + + +@dataclass +class TTSResult: + """Represents the result of a run of the TTS model on a batch. + + Args: + frames: list of long tensors with shape `[B, 1 + Q, 1]` representing + the audio and text tokens for each step. Note that acoustic delay is already corrected + at that point. + logged_text_tokens: for debugging, list of tuples `(predicted_tokens, next_input_token)`. + end_steps: gives the last valid step in `frames` for each item in the batch, or `None` if the + full text could not be fully synthesized within the provided budget. + all_consumption_times: for each item in the batch, a list of steps at which individual entries + (see `Entry`) were consumed as input to the model. + all_transcripts: for each item in the batch, a list of pairs `(word, step)` indicating + at which step the given `word` in the transcript should appear. Divide by the frame rate + to obtain a time stamp. + """ + frames: list[torch.Tensor] + logged_text_tokens: list[list[tuple[int, int]]] + end_steps: list[int | None] + all_consumption_times: list[list[int]] + all_transcripts: list[list[tuple[str, int]]] + + +@dataclass +class TTSModel: + """Wrapper around a multi-stream language model, a mimi codec, and a text tokenizer that + provides the functionality of a TTS model. As an end-user, you should use `from_checkpoint_info` + rather than trying to build a TTSModel directly. + + Args: + lm: trained delayed streams model. + mimi: codec to use. + tokenizer: text tokenizer to use. + machine: TTS state machine to use, which depend on how the model was trained. + delay_steps: delay between the text and audio in steps. + max_speakers: maximum number of speakers in the cross attention for this model. + temp: temperature (for both text and audio). + cfg_coef: classifier free guidance coefficient. Note that some models were trained with + CFG distillation, e.g. CFG should not be used at inference time. + final_padding: how many steps to sample past the last word. + n_q: how many audio codebooks (e.g. RVQ levels) to generate. Trade off between quality and speed. + max_gen_length: will stop generating after that many steps even if the text has not been fully consumed. + padding_bonus: additive bonus for the padding logits, positive value will lead to slower speech. + kwargs: other arguments for `moshi.models.lm.LMGen`. + + """ + + # the following params will be automatically set by `from_checkpoint_info` + lm: LMModel + mimi: MimiModel + tokenizer: SentencePieceProcessor + + voice_suffix: str + voice_repo: str + + machine: StateMachine + delay_steps: int + max_speakers: int = 5 + + # The following params can be overriden to customize generation. + temp: float = 0.6 + cfg_coef: float = 1.0 + final_padding: int = 4 + n_q: int = 32 + max_gen_length: int = 30000 + padding_bonus: float = 0. + + @staticmethod + def from_checkpoint_info(checkpoint_info: loaders.CheckpointInfo, + initial_padding: int = 2, + max_padding: int = 8, + voice_repo: str = DEFAULT_DSM_TTS_VOICE_REPO, + device: torch.device | str = 'cpu', + dtype: torch.dtype = torch.bfloat16, **kwargs) -> 'TTSModel': + assert checkpoint_info.raw_config is not None + model_id = checkpoint_info.raw_config['model_id'] + voice_suffix = f".{model_id['sig']}@{model_id['epoch']}.safetensors" + + mimi = checkpoint_info.get_mimi(device=device) + tokenizer = checkpoint_info.get_text_tokenizer() + lm = checkpoint_info.get_moshi(device=device, dtype=dtype) + + token_ids = TokenIds(lm.text_card + 1) + delay_steps = int(checkpoint_info.tts_config['audio_delay'] * mimi.frame_rate) + second_stream_ahead = checkpoint_info.tts_config.get('second_stream_ahead', 0) + + machine = StateMachine( + token_ids=token_ids, second_stream_ahead=second_stream_ahead, + max_padding=max_padding, initial_padding=initial_padding) + tts_model = TTSModel( + lm=lm, mimi=mimi, tokenizer=tokenizer, + voice_suffix=voice_suffix, voice_repo=voice_repo, + machine=machine, delay_steps=delay_steps, + **kwargs) + mimi.set_num_codebooks(tts_model.n_q) + if not tts_model.multi_speaker: + tts_model.voice_suffix = '' + return tts_model + + @cached_property + def valid_cfg_conditionings(self) -> set[float]: + valid_cfg_conditionings = set() + if self.lm.condition_provider is not None and 'cfg' in self.lm.condition_provider.conditioners: + cfg_conditioner = self.lm.condition_provider.conditioners['cfg'] + assert isinstance(cfg_conditioner, LUTConditioner) + assert cfg_conditioner.tokenizer.possible_values is not None + valid_cfg_conditionings = set(float(x) for x in cfg_conditioner.tokenizer.possible_values) + return valid_cfg_conditionings + + @cached_property + def multi_speaker(self) -> bool: + if self.lm.condition_provider is None: + return False + return 'speaker_wavs' in self.lm.condition_provider.conditioners + + def prepare_script(self, script: tp.Sequence[str], padding_between: int = 0) -> list[Entry]: + """Wrapper around `script_to_entries`.""" + return script_to_entries( + self.tokenizer, self.machine.token_ids, self.mimi.frame_rate, script, + multi_speaker=self.multi_speaker, padding_between=padding_between) + + @torch.no_grad() + def generate(self, all_entries: tp.Sequence[tp.Sequence[Entry]], + attributes: tp.Sequence[ConditionAttributes], + prefixes: list[torch.Tensor] | None = None, + cfg_is_no_prefix: bool = True, + cfg_is_no_text: bool = True, + on_frame: tp.Optional[tp.Callable[[torch.Tensor], None]] = None, + **kwargs + ) -> TTSResult: + """Synthesize text to audio. Returns a `TTSResult`. + + Args: + all_entries: list with one item per batch item, consisting of a list of `Entry`, + obtained from `prepare_script`. + attributes: list of `ConditionAttributes` for speaker conditioning. + prefixes: this should be the list of the lengths up until when to mask for the CFG. + cfg_is_no_prefix: if true, the null logits are computed with a masked prefix. + cfg_is_no_text: if true, the null logits are computed without the text. + on_frame: a callback triggered when a frame of mimi codes is available, the frame + is a view on a pre-allocated tensor so has to be copied if you want to keep it. + **kwargs: passed to `moshi.models.lm.LMGen`. + """ + + if self.cfg_coef != 1.0: + if self.valid_cfg_conditionings: + raise ValueError( + "This model does not support direct CFG, but was trained with " + "CFG distillation. Pass instead `cfg_coef` to `make_condition_attributes`.") + nulled = _make_null(attributes) + attributes = list(attributes) + nulled + + assert self.lm.condition_provider is not None + prepared = self.lm.condition_provider.prepare(attributes) + condition_tensors = self.lm.condition_provider(prepared) + + states = [] + for entries in all_entries: + state = self.machine.new_state(entries) + states.append(state) + + cfg_is_masked_until = None + text_prefixes = None + audio_prefixes = None + device = self.lm.device + if prefixes is not None: + assert len(all_entries) == len(prefixes), f"Not enough prefixes, expected {len(all_entries)}." + if cfg_is_no_prefix: + cfg_is_masked_until = [] + text_prefixes = [] + audio_prefixes = [] + for prefix in prefixes: + if cfg_is_masked_until is not None: + cfg_is_masked_until.append(prefix.shape[-1] + self.delay_steps) + K, _ = prefix.shape + assert K == self.lm.num_codebooks + text_prefixes.append(deque(prefix[0].cpu().tolist())) + delays = [d + self.delay_steps for d in self.lm.delays[self.lm.audio_offset:]] + delayed = _delayed(prefix[self.lm.audio_offset:], delays, self.machine.token_ids.ungenerated) + delayed = delayed.to(device) + audio_prefixes.append(deque(delayed.t())) + + def _on_text_logits_hook(text_logits): + if self.padding_bonus: + text_logits[..., self.machine.token_ids.pad] += self.padding_bonus + return text_logits + + def _on_audio_hook(audio_tokens): + audio_offset = self.lm.audio_offset + delays = self.lm.delays + ungenerated = self.machine.token_ids.ungenerated + for q in range(audio_tokens.shape[1]): + delay = delays[q + audio_offset] + if offset < delay + self.delay_steps: + audio_tokens[:, q] = self.machine.token_ids.zero + if audio_prefixes is not None: + for b, audio_prefix in enumerate(audio_prefixes): + if audio_prefix: + audio_codes = audio_prefix.popleft() + mask = audio_codes != ungenerated + audio_tokens[b] = torch.where(mask, audio_codes, audio_tokens[b]) + + def _on_text_hook(text_tokens): + tokens = text_tokens.tolist() + out_tokens = [] + for b, (token, state, logged) in enumerate(zip(tokens, states, logged_text_tokens)): + if text_prefixes is not None and text_prefixes[b]: + out_token = text_prefixes[b].popleft() + else: + out_token, _ = self.machine.process(offset, state, token) + out_tokens.append(out_token) + logged.append((token, out_token)) + text_tokens[:] = torch.tensor(out_tokens, dtype=torch.long, device=text_tokens.device) + + self.lm.dep_q = self.n_q + lm_gen = LMGen( + self.lm, temp=self.temp, temp_text=self.temp, + cfg_coef=self.cfg_coef, condition_tensors=condition_tensors, + on_text_logits_hook=_on_text_logits_hook, on_text_hook=_on_text_hook, on_audio_hook=_on_audio_hook, + cfg_is_masked_until=cfg_is_masked_until, cfg_is_no_text=cfg_is_no_text, + **kwargs) + + logged_text_tokens = [[] for _ in states] + frames: list[torch.Tensor] = [] + + with lm_gen.streaming(len(states)): + for offset in range(self.max_gen_length): + if all(state.end_step is not None for state in states): + max_end_step = max(state.end_step for state in states) + if offset >= max_end_step + self.delay_steps + self.final_padding: + break + missing = self.lm.n_q - self.lm.dep_q + input_tokens = torch.full((len(states), missing, 1), self.machine.token_ids.zero, + dtype=torch.long, device=self.lm.device) + frame = lm_gen.step(input_tokens) + if frame is not None: + frames.append(frame.clone()) + if on_frame is not None: + on_frame(frame) + return TTSResult( + frames, logged_text_tokens, + [state.end_step for state in states], + [state.consumption_times for state in states], + [state.transcript for state in states]) + + def get_voice_path(self, voice_name: str) -> Path: + """Returns a local path given a voice name, potentially fetching the voice + from a HuggingFace repository. To retrieve a voice from another repo, you can also use + the `hf://REPO/PATH` syntax. + """ + file = loaders.hf_get(voice_name + self.voice_suffix, self.voice_repo, + check_local_file_exists=True) + return Path(file) + + def make_condition_attributes( + self, voices: list[Path], cfg_coef: float | None = None) -> ConditionAttributes: + """Given a list of pre computed voice embeddings, returns a ConditionAttributes. + + Args: + voices: list of file paths to pre computed voice embeddings, see `get_voice_path`. + cfg_coef: for model trained with CFG distillation, value of the CFG + to use as conditioning. Typically, values from 1. to 4. are supported + with 0.5 increments. + """ + if voices: + voice_tensor = None + mask = None + for idx in range(5): + if idx < len(voices): + emb = load_file(voices[idx], device='cpu')['speaker_wavs'] + assert emb.dim() == 3 + if voice_tensor is None: + voice_tensor = torch.zeros(1, self.max_speakers, emb.shape[2], emb.shape[1]) + if mask is None: + mask = torch.zeros(1, self.max_speakers, emb.shape[2], dtype=torch.bool) + voice_tensor[:, idx, :, :] = emb.transpose(1, 2) + mask[:, idx, :] = True + assert voice_tensor is not None + assert mask is not None + voice_tensor = voice_tensor.view(1, -1, voice_tensor.shape[-1]) + mask = mask.view(1, -1) + tensors = { + 'speaker_wavs': TensorCondition(voice_tensor, mask) + } + else: + tensors = {} + text: dict[str, str | None] = {'control': 'ok'} + if cfg_coef is None: + text['cfg'] = None + else: + if cfg_coef in self.valid_cfg_conditionings: + text['cfg'] = format(cfg_coef, '.1f') + else: + valids = ", ".join(str(x) for x in self.valid_cfg_conditionings) + raise ValueError(f"Unsupported value for cfg_coef, valid values are {valids}.") + return ConditionAttributes(text=text, tensor=tensors) + + def get_prefix(self, audio_path: Path) -> torch.Tensor: + wav, _ = sphn.read(audio_path, sample_rate=self.mimi.sample_rate) + with torch.no_grad(): + prefix = self.mimi.encode(torch.from_numpy(wav).to(device=self.lm.device)[None])[0, :, :-2] + null_text = torch.full_like(prefix[:1], self.machine.token_ids.zero) + prefix = torch.cat([null_text, prefix], dim=0) + return prefix diff --git a/moshi_src/moshi/modules/__init__.py b/moshi_src/moshi/modules/__init__.py new file mode 100644 index 0000000..c24cbf9 --- /dev/null +++ b/moshi_src/moshi/modules/__init__.py @@ -0,0 +1,23 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +"""Modules used for building the models.""" + +# flake8: noqa +from .conv import ( + NormConv1d, + NormConvTranspose1d, + StreamingConv1d, + StreamingConvTranspose1d, + pad_for_conv1d, + pad1d, + unpad1d, +) +from .seanet import SEANetEncoder, SEANetDecoder +from .transformer import StreamingTransformer diff --git a/moshi_src/moshi/modules/conv.py b/moshi_src/moshi/modules/conv.py new file mode 100644 index 0000000..f36b5cd --- /dev/null +++ b/moshi_src/moshi/modules/conv.py @@ -0,0 +1,423 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +from dataclasses import dataclass +import itertools +import math +import typing as tp +import warnings + +import torch +from torch import nn +from torch.nn import functional as F +from torch.nn.utils import weight_norm + +from .streaming import StreamingModule, State + + +CONV_NORMALIZATIONS = frozenset(["none", "weight_norm"]) +M = tp.TypeVar('M', bound=nn.Module) + + +class TransposedLayerNorm(nn.Module): + """LayerNorm for [B, C, T] inputs.""" + + def __init__(self, **kwargs): + super().__init__() + self.layer_norm = nn.LayerNorm(**kwargs) + + def forward(self, x): + x = x.transpose(1, 2) + x = self.layer_norm(x) + return x.transpose(1, 2) + + +def apply_parametrization_norm(module: M, norm: str = "none") -> M: + assert norm in CONV_NORMALIZATIONS + if norm == "weight_norm": + return weight_norm(module) + else: + # We already check was in CONV_NORMALIZATION, so any other choice + # doesn't need reparametrization. + return module + + +def get_extra_padding_for_conv1d( + x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0 +) -> int: + """See `pad_for_conv1d`.""" + length = x.shape[-1] + n_frames = (length - kernel_size + padding_total) / stride + 1 + ideal_length = (math.ceil(n_frames) - 1) * stride + (kernel_size - padding_total) + return ideal_length - length + + +def pad_for_conv1d( + x: torch.Tensor, kernel_size: int, stride: int, padding_total: int = 0 +): + """Pad for a convolution to make sure that the last window is full. + Extra padding is added at the end. This is required to ensure that we can rebuild + an output of the same length, as otherwise, even with padding, some time steps + might get removed. + For instance, with total padding = 4, kernel size = 4, stride = 2: + 0 0 1 2 3 4 5 0 0 # (0s are padding) + 1 2 3 # (output frames of a convolution, last 0 is never used) + 0 0 1 2 3 4 5 0 # (output of tr. conv., but pos. 5 is going to get removed as padding) + 1 2 3 4 # once you removed padding, we are missing one time step ! + """ + extra_padding = get_extra_padding_for_conv1d(x, kernel_size, stride, padding_total) + return F.pad(x, (0, extra_padding)) + + +def pad1d( + x: torch.Tensor, + paddings: tp.Tuple[int, int], + mode: str = "constant", + value: float = 0.0, +): + """Tiny wrapper around F.pad, just to allow for reflect padding on small input. + If this is the case, we insert extra 0 padding to the right before the reflection happen. + """ + length = x.shape[-1] + padding_left, padding_right = paddings + assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right) + if mode == "reflect": + max_pad = max(padding_left, padding_right) + extra_pad = 0 + if length <= max_pad: + extra_pad = max_pad - length + 1 + x = F.pad(x, (0, extra_pad)) + padded = F.pad(x, paddings, mode, value) + end = padded.shape[-1] - extra_pad + return padded[..., :end] + else: + return F.pad(x, paddings, mode, value) + + +def unpad1d(x: torch.Tensor, paddings: tp.Tuple[int, int]): + """Remove padding from x, handling properly zero padding. Only for 1d!""" + padding_left, padding_right = paddings + assert padding_left >= 0 and padding_right >= 0, (padding_left, padding_right) + assert (padding_left + padding_right) <= x.shape[-1] + end = x.shape[-1] - padding_right + return x[..., padding_left:end] + + +class NormConv1d(nn.Module): + """Wrapper around Conv1d and normalization applied to this conv + to provide a uniform interface across normalization approaches. + """ + + def __init__( + self, + *args, + causal: bool = False, + norm: str = "none", + norm_kwargs: tp.Dict[str, tp.Any] = {}, + **kwargs, + ): + super().__init__() + self.conv = apply_parametrization_norm( + nn.Conv1d(*args, **kwargs), norm + ) + self.norm_type = norm + + def forward(self, x): + x = self.conv(x) + return x + + +class NormConvTranspose1d(nn.Module): + """Wrapper around ConvTranspose1d and normalization applied to this conv + to provide a uniform interface across normalization approaches. + """ + + def __init__( + self, + *args, + causal: bool = False, + norm: str = "none", + norm_kwargs: tp.Dict[str, tp.Any] = {}, + **kwargs, + ): + super().__init__() + self.convtr = apply_parametrization_norm( + nn.ConvTranspose1d(*args, **kwargs), norm + ) + self.norm_type = norm + + def forward(self, x): + x = self.convtr(x) + return x + + +@dataclass +class _StreamingConv1dState(State): + previous: torch.Tensor + first: torch.Tensor + + def reset(self, reset_mask: torch.Tensor): + super().reset(reset_mask) + self.previous[:] = torch.where(reset_mask.view(-1, 1, 1), torch.zeros_like(self.previous), self.previous) + self.first[:] = torch.where(reset_mask, torch.ones_like(self.first), self.first) + + +class StreamingConv1d(StreamingModule[_StreamingConv1dState]): + """Conv1d with some builtin handling of asymmetric or causal padding + and normalization. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int, + stride: int = 1, + dilation: int = 1, + groups: int = 1, + bias: bool = True, + causal: bool = False, + norm: str = "none", + norm_kwargs: tp.Dict[str, tp.Any] = {}, + pad_mode: str = "constant", + ): + super().__init__() + assert pad_mode in ['constant', 'replicate'], pad_mode + self.pad_mode = pad_mode + assert causal + # warn user on unusual setup between dilation and stride + if stride > 1 and dilation > 1: + warnings.warn( + "StreamingConv1d has been initialized with stride > 1 and dilation > 1" + f" (kernel_size={kernel_size} stride={stride}, dilation={dilation})." + ) + self.conv = NormConv1d( + in_channels, + out_channels, + kernel_size, + stride, + dilation=dilation, + groups=groups, + bias=bias, + causal=causal, + norm=norm, + norm_kwargs=norm_kwargs, + ) + + @property + def _stride(self) -> int: + return self.conv.conv.stride[0] + + @property + def _kernel_size(self) -> int: + return self.conv.conv.kernel_size[0] + + @property + def _effective_kernel_size(self) -> int: + dilation = self.conv.conv.dilation[0] + return ( + self._kernel_size - 1 + ) * dilation + 1 # effective kernel size with dilations + + @property + def _padding_total(self) -> int: + return self._effective_kernel_size - self._stride + + def _init_streaming_state(self, batch_size: int) -> _StreamingConv1dState: + stride = self._stride + # Effective kernel size accounting for dilation. + kernel = self._effective_kernel_size + param = next(iter(self.parameters())) + dtype = param.dtype + device = param.device + previous = torch.zeros(batch_size, self.conv.conv.in_channels, kernel - stride, + dtype=dtype, device=device) + first = torch.ones(batch_size, device=device, dtype=torch.bool) + return _StreamingConv1dState(batch_size, device, previous, first) + + def forward(self, x): + B, C, T = x.shape + S = self._stride + assert T > 0 and T % S == 0, "Steps must be multiple of stride" + state = self._streaming_state + if state is None: + state = self._init_streaming_state(B) + TP = state.previous.shape[-1] + if TP and self.pad_mode == 'replicate': + assert T >= TP, "Not enough content to pad streaming." + init = x[..., :1] + state.previous[:] = torch.where( + state.first.view(-1, 1, 1) & state.exec_mask.view(-1, 1, 1), + init, + state.previous) + if TP: + x = torch.cat([state.previous, x], dim=-1) + y = self.conv(x) + if TP: + state.previous[:] = torch.where( + state.exec_mask.view(-1, 1, 1), + x[..., -TP:], + state.previous) + if self.pad_mode == 'replicate': + state.first = torch.where( + state.exec_mask, + torch.zeros_like(state.first), + state.first, + ) + return y + + +@dataclass +class _StreamingConvTr1dState(State): + partial: torch.Tensor + + def reset(self, reset_mask: torch.Tensor): + super().reset(reset_mask) + self.partial[:] = torch.where( + reset_mask.view(-1, 1, 1), + torch.zeros_like(self.partial), + self.partial) + + +class StreamingConvTranspose1d(StreamingModule[_StreamingConvTr1dState]): + """ConvTranspose1d with some builtin handling of asymmetric or causal padding + and normalization. + """ + + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int, + stride: int = 1, + groups: int = 1, + bias: bool = True, + causal: bool = False, + norm: str = "none", + trim_right_ratio: float = 1.0, + norm_kwargs: tp.Dict[str, tp.Any] = {}, + ): + super().__init__() + assert trim_right_ratio == 1. + assert causal + self.convtr = NormConvTranspose1d( + in_channels, + out_channels, + kernel_size, + stride, + groups=groups, + bias=bias, + causal=causal, + norm=norm, + norm_kwargs=norm_kwargs, + ) + + @property + def _stride(self) -> int: + return self.convtr.convtr.stride[0] + + @property + def _kernel_size(self) -> int: + return self.convtr.convtr.kernel_size[0] + + def _init_streaming_state(self, batch_size: int) -> _StreamingConvTr1dState: + param = next(iter(self.parameters())) + dtype = param.dtype + device = param.device + K = self._kernel_size + S = self._stride + partial = torch.zeros(batch_size, self.convtr.convtr.out_channels, K - S, + device=device, dtype=dtype) + return _StreamingConvTr1dState(batch_size, device, partial) + + def forward(self, x): + B, C, T = x.shape + K = self._kernel_size + S = self._stride + state = self._streaming_state + + y = self.convtr(x) + if state is None: + y = unpad1d(y, (0, K - S)) + else: + PT = state.partial.shape[-1] + if PT > 0: + y[..., :PT] += state.partial + bias = self.convtr.convtr.bias + for_partial = y[..., -PT:] + if bias is not None: + for_partial -= bias[:, None] + state.partial[:] = torch.where( + state.exec_mask.view(-1, 1, 1), + for_partial, + state.partial) + y = y[..., :-PT] + return y + + +def test(): + torch.manual_seed(1234) + device = "cpu" + if torch.cuda.is_available(): + # Avoid the cuda optimizations that would take place on single precision + # floats for convolutions. + torch.backends.cudnn.enabled = True + torch.backends.cudnn.benchmark = False + torch.backends.cudnn.deterministic = True + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + device = "cuda:0" + + kernel_sizes = [1, 3, 4, 8, 15, 16] + strides = [1, 2, 3, 4, 5, 6, 7, 8, 9] + chin = 6 + chout = 12 + + for kernel, stride in itertools.product(kernel_sizes, strides): + if stride > kernel: + continue + conv = StreamingConv1d(chin, chout, kernel, stride, causal=True).to(device) + convtr = StreamingConvTranspose1d(chout, chin, kernel, stride, causal=True).to(device) + + for frames in [1, 4, 8, 32, 54, 65, 128]: + print(f"ksize {kernel} strides {stride} frames {frames}") + batch_size = 3 + length = frames * stride + x = torch.randn(batch_size, chin, length).to(device) + y = conv(x) + z = convtr(y) + for chunk_frames in [1, 2, 8]: + if frames % chunk_frames != 0: + continue + ys = [] + zs = [] + chunk_length = chunk_frames * stride + with conv.streaming(batch_size), convtr.streaming(batch_size): + for offset in range(0, length, chunk_length): + chunk = x[..., offset : offset + chunk_length] + ys.append(conv(chunk)) + zs.append(convtr(ys[-1])) + y_stream = torch.cat(ys, dim=-1) + z_stream = torch.cat(zs, dim=-1) + y = y[..., : y_stream.shape[-1]] + z = z[..., : z_stream.shape[-1]] + assert y.shape == y_stream.shape, (y.shape, y_stream.shape) + delta = (y_stream - y).norm() / y.norm() + assert delta <= 1e-6, delta + assert frames == y_stream.shape[-1], (frames, y_stream.shape) + + assert z.shape == z_stream.shape, (z.shape, z_stream.shape) + delta = (z_stream - z).norm() / z.norm() + assert delta <= 1e-6, (delta, (z_stream - z).abs().mean(dim=(0, 1))) + + +if __name__ == "__main__": + with torch.no_grad(): + test() diff --git a/moshi_src/moshi/modules/conv_test.py b/moshi_src/moshi/modules/conv_test.py new file mode 100644 index 0000000..8bd3228 --- /dev/null +++ b/moshi_src/moshi/modules/conv_test.py @@ -0,0 +1,157 @@ +import functools +import torch +import torch.nn as nn +import pytest + +from .conv import StreamingConv1d, StreamingConvTranspose1d + + +torch.backends.cudnn.enabled = False # Disable cuDNN for deterministic behavior and for numerical stability + + +CONV1D_DATA = [ + # batch_size, in_channels, out_channels, seq_len, kernel_size + pytest.param( + 3, 4, 5, 10, 6, + id='small conv1d test 1', + ), + pytest.param( + 4, 5, 6, 10, 7, + id='small conv1d test 2', + ), + pytest.param( + 5, 6, 7, 10, 2, + id='small conv1d test 3', + ), + pytest.param( + 1, 512, 512, 256, 7, + id='large conv1d test 1', + ), +] + +CONV1D_TRANSPOSE_DATA = [ + # batch_size, in_channels, out_channels, seq_len, kernel_size, stride + pytest.param( + 3, 4, 5, 10, 6, 1, + id='small conv1d transpose test 1', + ), + pytest.param( + 4, 5, 6, 10, 7, 2, + id='small conv1d transpose test 2', + ), + pytest.param( + 5, 6, 7, 10, 4, 3, + id='small conv1d transpose test 3', + ), + pytest.param( + 1, 512, 512, 256, 7, 2, + id='large conv1d transpose test 1', + ), +] + + +def _init_weights(module, generator=None): + for name, param in module.named_parameters(): + if "weight" in name: + nn.init.xavier_uniform_(param, generator=generator) + elif "bias" in name: + nn.init.constant_(param, 0.0) + else: + nn.init.xavier_uniform_(param, generator=generator) + + +@pytest.mark.parametrize("batch_size, in_channels, out_channels, seq_len, kernel_size", CONV1D_DATA) +def test_conv1d(batch_size, in_channels, out_channels, seq_len, kernel_size): + """Test that StreamingConv1d() calls are causal. Having new inputs does not change the previous output.""" + assert seq_len > kernel_size + + layer = StreamingConv1d(in_channels, out_channels, kernel_size, causal=True, norm="none", pad_mode="constant") + + generator = torch.Generator() + generator = generator.manual_seed(41) + layer.apply(functools.partial(_init_weights, generator=generator)) + + shape = (batch_size, in_channels, seq_len,) + input_hidden_states = torch.rand(shape) + + expected_output = layer(input_hidden_states) + + for end_index in range(kernel_size, seq_len + 1): + actual_output = layer(input_hidden_states[..., :end_index]) + torch.testing.assert_close(actual_output, expected_output[..., :actual_output.shape[-1]], + msg=lambda original_msg: f"Failed at end_index={end_index}: \n{original_msg}") + + +@pytest.mark.parametrize("batch_size, in_channels, out_channels, seq_len, kernel_size", CONV1D_DATA) +def test_conv1d_streaming(batch_size, in_channels, out_channels, seq_len, kernel_size): + """Test that StreamingConv1d() streaming works as expected.""" + assert seq_len > kernel_size + + layer = StreamingConv1d(in_channels, out_channels, kernel_size, causal=True, norm="none", pad_mode="constant") + + generator = torch.Generator() + generator = generator.manual_seed(41) + layer.apply(functools.partial(_init_weights, generator=generator)) + + shape = (batch_size, in_channels, seq_len,) + input_hidden_states = torch.rand(shape) + expected_output = layer(input_hidden_states) + + start_index = 0 + actual_outputs = [] + with layer.streaming(batch_size=batch_size): + for end_index in range(kernel_size, seq_len + 1): + actual_output = layer(input_hidden_states[..., start_index:end_index]) + start_index = end_index + actual_outputs.append(actual_output) + actual_outputs = torch.cat(actual_outputs, dim=-1) + + torch.testing.assert_close(actual_outputs, expected_output) + + +@pytest.mark.parametrize("batch_size, in_channels, out_channels, seq_len, kernel_size, stride", CONV1D_TRANSPOSE_DATA) +def test_conv1d_transpose(batch_size, in_channels, out_channels, seq_len, kernel_size, stride): + """Test that StreamingConvTranspose1d() calls are causal. Having new inputs does not change the previous output.""" + assert seq_len > kernel_size + + layer = StreamingConvTranspose1d(in_channels, out_channels, kernel_size, stride, causal=True, norm="none") + + generator = torch.Generator() + generator = generator.manual_seed(41) + layer.apply(functools.partial(_init_weights, generator=generator)) + + shape = (batch_size, in_channels, seq_len,) + input_hidden_states = torch.rand(shape) + expected_output = layer(input_hidden_states) + + for end_index in range(kernel_size, seq_len + 1): + actual_output = layer(input_hidden_states[..., :end_index]) + torch.testing.assert_close(actual_output, expected_output[..., :actual_output.shape[-1]], + msg=lambda original_msg: f"Failed at end_index={end_index}: \n{original_msg}") + + +@pytest.mark.parametrize("batch_size, in_channels, out_channels, seq_len, kernel_size, stride", CONV1D_TRANSPOSE_DATA) +def test_conv1d_transpose_streaming(batch_size, in_channels, out_channels, seq_len, kernel_size, stride): + """Test that StreamingConvTranspose1d() streaming works as expected.""" + assert seq_len > kernel_size + + layer = StreamingConvTranspose1d(in_channels, out_channels, kernel_size, stride, causal=True, norm="none") + + generator = torch.Generator() + generator = generator.manual_seed(41) + layer.apply(functools.partial(_init_weights, generator=generator)) + + shape = (batch_size, in_channels, seq_len,) + input_hidden_states = torch.rand(shape) + expected_output = layer(input_hidden_states) + + start_index = 0 + actual_outputs = [] + with layer.streaming(batch_size=batch_size): + for end_index in range(kernel_size, seq_len + 1): + actual_output = layer(input_hidden_states[..., start_index:end_index]) + start_index = end_index + actual_outputs.append(actual_output) + actual_outputs = torch.cat(actual_outputs, dim=-1) + + torch.testing.assert_close(actual_outputs, expected_output) diff --git a/moshi_src/moshi/modules/gating.py b/moshi_src/moshi/modules/gating.py new file mode 100644 index 0000000..bd708ad --- /dev/null +++ b/moshi_src/moshi/modules/gating.py @@ -0,0 +1,115 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +from contextlib import ExitStack +import torch +from torch import nn +from torch.nn import functional as F + +from ..utils.compile import torch_compile_lazy, no_compile + + + +def gating_forward_kernel( + weight_in: torch.Tensor, weight_out: torch.Tensor, activation, x: torch.Tensor +): + x = F.linear(x, weight_in) + B, T, _ = x.shape + x = x.view(B, T, 2, -1) + x = activation(x[..., 0, :]) * x[..., 1, :] + x = F.linear(x, weight_out) + return x + + +def gating_forward_generic( + linear_in: nn.Module, + linear_out: nn.Module, + activation, + x: torch.Tensor +): + x = linear_in(x) + B, T, _ = x.shape + x = x.view(B, T, 2, -1) + x = activation(x[..., 0, :]) * x[..., 1, :] + x = linear_out(x) + return x + + +class ActivationGating(nn.Module): + """ + Gating FFN layer, using the given activation. + Args: + dim (int): dimension of the input and output of the transformer. + activation (any callable Tensor to Tensor): activation function to use. + **factory_kwargs: other kwargs passed to the linear layer, in particular device and dtype. + """ + + _fsdp_final = True + + def __init__(self, dim: int, dim_feedforward: int, activation, quantized: bool = False, **factory_kwargs): + super().__init__() + # We should have 8 d^2 param, instead we will have + # 2 * h * d + h * d = 3 h * d = 8 d^2 + # so h = 8 d / 3 but following Hervé's advice we use 21 / 8 as an approx. + if dim_feedforward == 4 * dim: + hidden = (21 * dim) // 8 + else: + hidden = (2 * dim_feedforward) // 3 + + self.linear_in = nn.Linear(dim, 2 * hidden, bias=False, **factory_kwargs) + self.linear_out = nn.Linear(hidden, dim, bias=False, **factory_kwargs) + + # We try to follow the default PyTorch MHA convention, to easily compare results. + + self.activation = activation + + def forward(self, x: torch.Tensor): + if isinstance(self.linear_in, nn.Linear): + assert isinstance(self.linear_out, nn.Linear) + with ExitStack() as stack: + if self.training: + stack.enter_context(no_compile()) + return gating_forward_kernel( + self.linear_in.weight, self.linear_out.weight, self.activation, x + ) + else: + return gating_forward_generic( + self.linear_in, + self.linear_out, + self.activation, + x + ) + + +def _get_activation(name: str): + if name in ["sigmoid", "tanh", "relu"]: + return getattr(torch, name) + elif name in ["leaky_relu", "elu", "gelu", "silu", "mish", "softsign"]: + return getattr(torch.nn.functional, name) + elif name == "identity": + return torch.nn.Identity() + else: + raise ValueError(f"Unknown activation {name}") + + +def _make_gating( + name: str, dim: int, dim_feedforward: int, + **factory_kwargs +) -> nn.Module: + return ActivationGating( + dim, dim_feedforward, _get_activation(name), **factory_kwargs + ) + + +def make_gating( + name: str, dim: int, dim_feedforward: int, **factory_kwargs +) -> nn.Module: + gating = _make_gating(name, dim, dim_feedforward, **factory_kwargs) + if isinstance(gating.linear_in, nn.Linear): + max_params = 2 * dim * dim_feedforward + params = sum(p.numel() for p in gating.parameters()) + assert ( + params <= max_params + ), f"{name} gating has {params} params, max is {max_params}" + return gating diff --git a/moshi_src/moshi/modules/lora.py b/moshi_src/moshi/modules/lora.py new file mode 100644 index 0000000..b4731bc --- /dev/null +++ b/moshi_src/moshi/modules/lora.py @@ -0,0 +1,122 @@ +import torch +import torch.nn as nn + + +def replace_all_linear_with_lora(module, rank: int, scaling: float, device=None, dtype=None): + """ Recursively replace all Linear layers with LoRALinear layers.""" + for name, child in module.named_children(): + if isinstance(child, nn.Linear): + if device is None: + this_device = child.weight.device + else: + this_device = device + if dtype is None: + this_dtype = child.weight.dtype + else: + this_dtype = dtype + lora = LoRALinear(child.in_features, child.out_features, + rank, scaling, device=this_device, dtype=this_dtype) + lora.frozen_W = child + setattr(module, name, lora) + else: + replace_all_linear_with_lora(child, rank, scaling, device=device, dtype=dtype) + + +def replace_lora_with_linear(module): + """Recursively replace all LoRALinear layers with Linear layers.""" + for name, child in module.named_children(): + if isinstance(child, LoRALinear): + # Compute merged weights: W' = W + scaling * B @ A + merged_weight = child.frozen_W.weight.data + \ + child.scaling * (child.lora_B.weight @ child.lora_A.weight) + # Create a standard Linear layer with the same in/out features + new_linear = nn.Linear(child.frozen_W.in_features, + child.frozen_W.out_features, bias=False, + device=torch.device('meta'), + dtype=merged_weight.dtype) + new_linear.weight = nn.Parameter( + merged_weight, requires_grad=merged_weight.requires_grad) # Transfer merged weights + setattr(module, name, new_linear) # Replace the module + else: + replace_lora_with_linear(child) # Recursively process submodules + + +class LoRALinear(nn.Module): + """ + Implementation of: + - LoRA: https://arxiv.org/abs/2106.09685 + + Notes: + - Freezing is handled at the network level, not the layer level. + - Scaling factor controls relative importance of LoRA skip + connection versus original frozen weight. General guidance is + to keep it to 2.0 and sweep over learning rate when changing + the rank. + """ + + def __init__( + self, + in_features: int, + out_features: int, + rank: int, + scaling: float, + bias: bool = False, + device: torch.device | None = None, + dtype: torch.dtype = torch.bfloat16, + ): + super().__init__() + + self.in_features = in_features + self.out_features = out_features + assert not bias + self.bias = bias + self.rank = rank + self.scaling = scaling + + self.lora_A = nn.Linear( + self.in_features, + self.rank, + bias=self.bias, + device=device, + dtype=dtype, + ) + self.lora_B = nn.Linear( + self.rank, + self.out_features, + bias=self.bias, + device=device, + dtype=dtype, + ) + + self.frozen_W = nn.Linear(self.in_features, + self.out_features, + bias=self.bias, + device=device, + dtype=dtype) + + self._register_load_state_dict_pre_hook(LoRALinear._load_hook, with_module=True) + + def merge_weight(self): + with torch.no_grad(): + down_weight = self.lora_A.weight + up_weight = self.lora_B.weight + + weight = up_weight.mm(down_weight) * self.scaling + + weight += self.frozen_W.weight + return weight + + @staticmethod + def _load_hook(module, state_dict, prefix, *_): + key_name = prefix + "weight" + if key_name in state_dict: + w_ref = state_dict.pop(key_name) + state_dict[prefix + 'frozen_W.weight'] = w_ref + + def forward(self, x: torch.Tensor): + lora = self.lora_B(self.lora_A(x)) + return self.frozen_W(x) + lora * self.scaling + + def __repr__(self) -> str: + return "{}Linear(in_features={}, out_features={}, r={})".format( + "LoRA", self.in_features, self.out_features, self.rank) diff --git a/moshi_src/moshi/modules/resample.py b/moshi_src/moshi/modules/resample.py new file mode 100644 index 0000000..e260796 --- /dev/null +++ b/moshi_src/moshi/modules/resample.py @@ -0,0 +1,119 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import typing as tp + +from einops import rearrange +import torch +from torch import nn + +from .conv import StreamingConv1d, StreamingConvTranspose1d + + +class ConvDownsample1d(nn.Module): + """ + Downsampling by some integer amount `stride` using convolutions + with a kernel size of twice the stride. + If `causal` is True, the output uses a causal convolution. + """ + + def __init__( + self, + stride: int, + dimension: tp.Optional[int] = None, + causal: bool = False, + learnt: bool = False, + channel_wise: bool = False, + ): + super().__init__() + self.learnt = learnt + self.channel_wise = channel_wise + groups = 1 + if learnt: + assert dimension is not None, "Dimension required for learnt convolutions." + in_channels = dimension + out_channels = dimension + if channel_wise: + groups = dimension + else: + in_channels = 1 + out_channels = 1 + + self.conv = StreamingConv1d( + in_channels, + out_channels, + kernel_size=2 * stride, + stride=stride, + causal=causal, + groups=groups, + bias=False, + pad_mode="replicate", + ) + if not learnt: + actual_conv = self.conv.conv.conv + actual_conv.weight.requires_grad_(False) + actual_conv.weight.data.fill_(1.0 / (2 * stride)) + + def forward(self, x: torch.Tensor): + batch_size = len(x) + if not self.learnt: + x = rearrange(x, "b c t -> (b c) () t") + y = self.conv(x) + if not self.learnt: + y = rearrange(y, "(b c) () t -> b c t", b=batch_size) + return y + + +class ConvTrUpsample1d(nn.Module): + """ + Upsample by some integer amount `stride` using transposed convolutions. + """ + + def __init__( + self, + stride: int, + dimension: tp.Optional[int] = None, + causal: bool = False, + learnt: bool = False, + channel_wise: bool = False, + ): + super().__init__() + self.learnt = learnt + self.channel_wise = channel_wise + groups = 1 + if learnt: + assert dimension is not None, "Dimension required for learnt convolutions." + in_channels = dimension + out_channels = dimension + if channel_wise: + groups = dimension + else: + in_channels = 1 + out_channels = 1 + + self.convtr = StreamingConvTranspose1d( + in_channels, + out_channels, + kernel_size=2 * stride, + stride=stride, + causal=causal, + groups=groups, + bias=False, + ) + if not learnt: + actual_convtr = self.convtr.convtr.convtr + actual_convtr.weight.requires_grad_(False) + actual_convtr.weight.data.fill_(1.0) + + def forward(self, x: torch.Tensor): + batch_size = len(x) + if not self.learnt: + x = rearrange(x, "b c t -> (b c) () t") + y = self.convtr(x) + if not self.learnt: + x_for_normalization = torch.ones_like(x[:1]) + normalization = self.convtr(x_for_normalization) + y = y / normalization + y = rearrange(y, "(b c) () t -> b c t", b=batch_size) + return y diff --git a/moshi_src/moshi/modules/rope.py b/moshi_src/moshi/modules/rope.py new file mode 100644 index 0000000..0965501 --- /dev/null +++ b/moshi_src/moshi/modules/rope.py @@ -0,0 +1,90 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +from torch import nn +import math +import torch +from ..utils.compile import torch_compile_lazy + + + +def apply_rope( + q: torch.Tensor, + k: torch.Tensor, + offset: torch.Tensor, + max_period: float = 10_000, + time_before_heads: bool = False, +): + """ + Args: + q (torch.Tensor): queries, shape `[B, T, H, D]`. + k (torch.Tensor): keys, shape `[B, T, H, D]`. + offset (int): current offset, e.g. when streaming. + max_period (float): maximum period for the cos and sin. + time_before_heads (bool): if True, expected [B, T, H, D], else [B, H, T ,D] + """ + + if time_before_heads: + B, T, H, D = q.shape + else: + B, H, T, D = q.shape + assert k.shape == q.shape + assert D > 0 + assert D % 2 == 0 + assert max_period > 0 + + ds = torch.arange(D // 2, device=q.device, dtype=torch.float32) + freqs = torch.exp(ds * (-math.log(max_period) * 2 / D)) + ts = offset.float().view(-1, 1) + torch.arange(T, device=q.device, dtype=torch.float32) + if time_before_heads: + ts = ts.view(B, -1, 1, 1) + else: + ts = ts.view(B, 1, -1, 1) + + dims = q.shape[:-1] + q = q.view(*dims, D // 2, 2) + k = k.view(*dims, D // 2, 2) + + # convention is `r` suffix is real part, `i` is imaginary. + qr = q[..., 0].float() + qi = q[..., 1].float() + + kr = k[..., 0].float() + ki = k[..., 1].float() + + rotr = torch.cos(freqs * ts) + roti = torch.sin(freqs * ts) + qor = qr * rotr - qi * roti + qoi = qr * roti + qi * rotr + + kor = kr * rotr - ki * roti + koi = kr * roti + ki * rotr + + dtype = q.dtype + qo = torch.stack([qor.to(dtype), qoi.to(dtype)], dim=-1) + ko = torch.stack([kor.to(dtype), koi.to(dtype)], dim=-1) + + return qo.view(*dims, D), ko.view(*dims, D) + + +class RotaryEmbedding(nn.Module): + """Rotary positional embedding (RoPE) from [Su et al 2022](https://arxiv.org/abs/2104.09864). + + Args: + max_period (float): Maximum period of the rotation frequencies. + """ + + def __init__(self, max_period: float = 10000.0): + super().__init__() + self.max_period = max_period + + def forward( + self, + q: torch.Tensor, + k: torch.Tensor, + offset: torch.Tensor, + time_before_heads: bool = False, + ): + """Apply rope rotation to query or key tensor.""" + return apply_rope(q, k, offset, self.max_period, time_before_heads) diff --git a/moshi_src/moshi/modules/seanet.py b/moshi_src/moshi/modules/seanet.py new file mode 100644 index 0000000..0129190 --- /dev/null +++ b/moshi_src/moshi/modules/seanet.py @@ -0,0 +1,392 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import typing as tp + +import numpy as np +import torch.nn as nn + +from .conv import StreamingConv1d, StreamingConvTranspose1d +from .streaming import StreamingContainer + + +class SEANetResnetBlock(StreamingContainer): + """Residual block from SEANet model. + + Args: + dim (int): Dimension of the input/output. + kernel_sizes (list): List of kernel sizes for the convolutions. + dilations (list): List of dilations for the convolutions. + activation (str): Activation function. + activation_params (dict): Parameters to provide to the activation function. + norm (str): Normalization method. + norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution. + causal (bool): Whether to use fully causal convolution. + pad_mode (str): Padding mode for the convolutions. + compress (int): Reduced dimensionality in residual branches (from Demucs v3). + true_skip (bool): Whether to use true skip connection or a simple + (streamable) convolution as the skip connection. + """ + + def __init__( + self, + dim: int, + kernel_sizes: tp.List[int] = [3, 1], + dilations: tp.List[int] = [1, 1], + activation: str = "ELU", + activation_params: dict = {"alpha": 1.0}, + norm: str = "none", + norm_params: tp.Dict[str, tp.Any] = {}, + causal: bool = False, + pad_mode: str = "reflect", + compress: int = 2, + true_skip: bool = True, + ): + super().__init__() + assert len(kernel_sizes) == len( + dilations + ), "Number of kernel sizes should match number of dilations" + act = getattr(nn, activation) + hidden = dim // compress + block = [] + for i, (kernel_size, dilation) in enumerate(zip(kernel_sizes, dilations)): + in_chs = dim if i == 0 else hidden + out_chs = dim if i == len(kernel_sizes) - 1 else hidden + block += [ + act(**activation_params), + StreamingConv1d( + in_chs, + out_chs, + kernel_size=kernel_size, + dilation=dilation, + norm=norm, + norm_kwargs=norm_params, + causal=causal, + pad_mode=pad_mode, + ), + ] + self.block = nn.Sequential(*block) + self.shortcut: nn.Module + if true_skip: + self.shortcut = nn.Identity() + else: + self.shortcut = StreamingConv1d( + dim, + dim, + kernel_size=1, + norm=norm, + norm_kwargs=norm_params, + causal=causal, + pad_mode=pad_mode, + ) + + def forward(self, x): + u, v = self.shortcut(x), self.block(x) + assert u.shape == v.shape, (u.shape, v.shape, x.shape) + return u + v + + +class SEANetEncoder(StreamingContainer): + """SEANet encoder. + + Args: + channels (int): Audio channels. + dimension (int): Intermediate representation dimension. + n_filters (int): Base width for the model. + n_residual_layers (int): nb of residual layers. + ratios (Sequence[int]): kernel size and stride ratios. The encoder uses downsampling ratios instead of + upsampling ratios, hence it will use the ratios in the reverse order to the ones specified here + that must match the decoder order. We use the decoder order as some models may only employ the decoder. + activation (str): Activation function. + activation_params (dict): Parameters to provide to the activation function. + norm (str): Normalization method. + norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution. + kernel_size (int): Kernel size for the initial convolution. + last_kernel_size (int): Kernel size for the initial convolution. + residual_kernel_size (int): Kernel size for the residual layers. + dilation_base (int): How much to increase the dilation with each layer. + causal (bool): Whether to use fully causal convolution. + pad_mode (str): Padding mode for the convolutions. + true_skip (bool): Whether to use true skip connection or a simple + (streamable) convolution as the skip connection in the residual network blocks. + compress (int): Reduced dimensionality in residual branches (from Demucs v3). + disable_norm_outer_blocks (int): Number of blocks for which we don't apply norm. + For the encoder, it corresponds to the N first blocks. + mask_fn (nn.Module): Optional mask function to apply after convolution layers. + mask_position (int): Position of the mask function, with mask_position == 0 for the first convolution layer, + mask_position == 1 for the first conv block, etc. + """ + + def __init__( + self, + channels: int = 1, + dimension: int = 128, + n_filters: int = 32, + n_residual_layers: int = 3, + ratios: tp.List[int] = [8, 5, 4, 2], + activation: str = "ELU", + activation_params: dict = {"alpha": 1.0}, + norm: str = "none", + norm_params: tp.Dict[str, tp.Any] = {}, + kernel_size: int = 7, + last_kernel_size: int = 7, + residual_kernel_size: int = 3, + dilation_base: int = 2, + causal: bool = False, + pad_mode: str = "reflect", + true_skip: bool = True, + compress: int = 2, + disable_norm_outer_blocks: int = 0, + mask_fn: tp.Optional[nn.Module] = None, + mask_position: tp.Optional[int] = None, + ): + super().__init__() + self.channels = channels + self.dimension = dimension + self.n_filters = n_filters + self.ratios = list(reversed(ratios)) + del ratios + self.n_residual_layers = n_residual_layers + self.hop_length = int(np.prod(self.ratios)) + self.n_blocks = len(self.ratios) + 2 # first and last conv + residual blocks + self.disable_norm_outer_blocks = disable_norm_outer_blocks + assert ( + self.disable_norm_outer_blocks >= 0 and self.disable_norm_outer_blocks <= self.n_blocks + ), ( + "Number of blocks for which to disable norm is invalid." + "It should be lower or equal to the actual number of blocks in the network and greater or equal to 0." + ) + + act = getattr(nn, activation) + mult = 1 + model: tp.List[nn.Module] = [ + StreamingConv1d( + channels, + mult * n_filters, + kernel_size, + norm="none" if self.disable_norm_outer_blocks >= 1 else norm, + norm_kwargs=norm_params, + causal=causal, + pad_mode=pad_mode, + ) + ] + if mask_fn is not None and mask_position == 0: + model += [mask_fn] + # Downsample to raw audio scale + for i, ratio in enumerate(self.ratios): + block_norm = "none" if self.disable_norm_outer_blocks >= i + 2 else norm + # Add residual layers + for j in range(n_residual_layers): + model += [ + SEANetResnetBlock( + mult * n_filters, + kernel_sizes=[residual_kernel_size, 1], + dilations=[dilation_base**j, 1], + norm=block_norm, + norm_params=norm_params, + activation=activation, + activation_params=activation_params, + causal=causal, + pad_mode=pad_mode, + compress=compress, + true_skip=true_skip, + ) + ] + + # Add downsampling layers + model += [ + act(**activation_params), + StreamingConv1d( + mult * n_filters, + mult * n_filters * 2, + kernel_size=ratio * 2, + stride=ratio, + norm=block_norm, + norm_kwargs=norm_params, + causal=causal, + pad_mode=pad_mode, + ), + ] + mult *= 2 + if mask_fn is not None and mask_position == i + 1: + model += [mask_fn] + + model += [ + act(**activation_params), + StreamingConv1d( + mult * n_filters, + dimension, + last_kernel_size, + norm=( + "none" if self.disable_norm_outer_blocks == self.n_blocks else norm + ), + norm_kwargs=norm_params, + causal=causal, + pad_mode=pad_mode, + ), + ] + + self.model = nn.Sequential(*model) + + def forward(self, x): + return self.model(x) + + +class SEANetDecoder(StreamingContainer): + """SEANet decoder. + + Args: + channels (int): Audio channels. + dimension (int): Intermediate representation dimension. + n_filters (int): Base width for the model. + n_residual_layers (int): nb of residual layers. + ratios (Sequence[int]): kernel size and stride ratios. + activation (str): Activation function. + activation_params (dict): Parameters to provide to the activation function. + final_activation (str): Final activation function after all convolutions. + final_activation_params (dict): Parameters to provide to the activation function. + norm (str): Normalization method. + norm_params (dict): Parameters to provide to the underlying normalization used along with the convolution. + kernel_size (int): Kernel size for the initial convolution. + last_kernel_size (int): Kernel size for the initial convolution. + residual_kernel_size (int): Kernel size for the residual layers. + dilation_base (int): How much to increase the dilation with each layer. + causal (bool): Whether to use fully causal convolution. + pad_mode (str): Padding mode for the convolutions. + true_skip (bool): Whether to use true skip connection or a simple. + (streamable) convolution as the skip connection in the residual network blocks. + compress (int): Reduced dimensionality in residual branches (from Demucs v3). + disable_norm_outer_blocks (int): Number of blocks for which we don't apply norm. + For the decoder, it corresponds to the N last blocks. + trim_right_ratio (float): Ratio for trimming at the right of the transposed convolution under the causal setup. + If equal to 1.0, it means that all the trimming is done at the right. + """ + + def __init__( + self, + channels: int = 1, + dimension: int = 128, + n_filters: int = 32, + n_residual_layers: int = 3, + ratios: tp.List[int] = [8, 5, 4, 2], + activation: str = "ELU", + activation_params: dict = {"alpha": 1.0}, + final_activation: tp.Optional[str] = None, + final_activation_params: tp.Optional[dict] = None, + norm: str = "none", + norm_params: tp.Dict[str, tp.Any] = {}, + kernel_size: int = 7, + last_kernel_size: int = 7, + residual_kernel_size: int = 3, + dilation_base: int = 2, + causal: bool = False, + pad_mode: str = "reflect", + true_skip: bool = True, + compress: int = 2, + disable_norm_outer_blocks: int = 0, + trim_right_ratio: float = 1.0, + ): + super().__init__() + self.dimension = dimension + self.channels = channels + self.n_filters = n_filters + self.ratios = ratios + del ratios + self.n_residual_layers = n_residual_layers + self.hop_length = int(np.prod(self.ratios)) + self.n_blocks = len(self.ratios) + 2 # first and last conv + residual blocks + self.disable_norm_outer_blocks = disable_norm_outer_blocks + assert ( + self.disable_norm_outer_blocks >= 0 and self.disable_norm_outer_blocks <= self.n_blocks + ), ( + "Number of blocks for which to disable norm is invalid." + "It should be lower or equal to the actual number of blocks in the network and greater or equal to 0." + ) + + act = getattr(nn, activation) + mult = int(2 ** len(self.ratios)) + model: tp.List[nn.Module] = [ + StreamingConv1d( + dimension, + mult * n_filters, + kernel_size, + norm=( + "none" if self.disable_norm_outer_blocks == self.n_blocks else norm + ), + norm_kwargs=norm_params, + causal=causal, + pad_mode=pad_mode, + ) + ] + + # Upsample to raw audio scale + for i, ratio in enumerate(self.ratios): + block_norm = ( + "none" + if self.disable_norm_outer_blocks >= self.n_blocks - (i + 1) + else norm + ) + # Add upsampling layers + model += [ + act(**activation_params), + StreamingConvTranspose1d( + mult * n_filters, + mult * n_filters // 2, + kernel_size=ratio * 2, + stride=ratio, + norm=block_norm, + norm_kwargs=norm_params, + causal=causal, + trim_right_ratio=trim_right_ratio, + ), + ] + # Add residual layers + for j in range(n_residual_layers): + model += [ + SEANetResnetBlock( + mult * n_filters // 2, + kernel_sizes=[residual_kernel_size, 1], + dilations=[dilation_base**j, 1], + activation=activation, + activation_params=activation_params, + norm=block_norm, + norm_params=norm_params, + causal=causal, + pad_mode=pad_mode, + compress=compress, + true_skip=true_skip, + ) + ] + + mult //= 2 + + # Add final layers + model += [ + act(**activation_params), + StreamingConv1d( + n_filters, + channels, + last_kernel_size, + norm="none" if self.disable_norm_outer_blocks >= 1 else norm, + norm_kwargs=norm_params, + causal=causal, + pad_mode=pad_mode, + ), + ] + # Add optional final activation to decoder (eg. tanh) + if final_activation is not None: + final_act = getattr(nn, final_activation) + final_activation_params = final_activation_params or {} + model += [final_act(**final_activation_params)] + self.model = nn.Sequential(*model) + + def forward(self, z): + y = self.model(z) + return y diff --git a/moshi_src/moshi/modules/seanet_test.py b/moshi_src/moshi/modules/seanet_test.py new file mode 100644 index 0000000..ce8c321 --- /dev/null +++ b/moshi_src/moshi/modules/seanet_test.py @@ -0,0 +1,187 @@ +import functools +import torch +import torch.nn as nn +import pytest + +from .seanet import SEANetResnetBlock, SEANetDecoder + + +torch.backends.cudnn.enabled = False # Disable cuDNN for deterministic behavior and for numerical stability + + +SEANET_RESNET_DATA = [ + # batch_size, dim, res_layer_index, seq_len, kernel_size + pytest.param( + 3, 4, 1, 10, 6, + id='small resnet test 1', + ), + pytest.param( + 4, 5, 2, 10, 7, + id='small resnet test 2', + ), + pytest.param( + 5, 6, 4, 10, 2, + id='small resnet test 3', + ), + pytest.param( + 1, 512, 2, 256, 7, + id='large resnet test 1', + ), +] +NUM_TIMESTEPS_DATA = [ + pytest.param( + 1, + id='length 1', + ), + pytest.param( + 2, + id='length 2', + ), + pytest.param( + 10, + id='length 10', + ), + pytest.param( + 100, + id='length 100', + ), +] + +SEANET_KWARGS_DATA = [ + pytest.param( + { + "channels": 1, + "dimension": 8, + "causal": True, + "n_filters": 2, + "n_residual_layers": 1, + "activation": "ELU", + "compress": 2, + "dilation_base": 2, + "disable_norm_outer_blocks": 0, + "kernel_size": 7, + "residual_kernel_size": 3, + "last_kernel_size": 3, + # We train using weight_norm but then the weights are pre-processed for inference so + # that we can use a normal convolution. + "norm": "none", + "pad_mode": "constant", + "ratios": [5], + "true_skip": True, + }, + id='Tiny SEANet', + ), + + pytest.param( + { + "channels": 1, + "dimension": 512, + "causal": True, + "n_filters": 64, + "n_residual_layers": 1, + "activation": "ELU", + "compress": 2, + "dilation_base": 2, + "disable_norm_outer_blocks": 0, + "kernel_size": 7, + "residual_kernel_size": 3, + "last_kernel_size": 3, + # We train using weight_norm but then the weights are pre-processed for inference so + # that we can use a normal convolution. + "norm": "none", + "pad_mode": "constant", + "ratios": [8, 6, 5, 4], + "true_skip": True, + }, + id='Large SEANet', + ), +] + + +def _init_weights(module, generator=None): + for name, param in module.named_parameters(): + if "weight" in name: + nn.init.xavier_uniform_(param, generator=generator) + elif "bias" in name: + nn.init.constant_(param, 0.0) + else: + nn.init.xavier_uniform_(param, generator=generator) + + +@pytest.mark.parametrize("batch_size, dim, res_layer_index, seq_len, kernel_size", SEANET_RESNET_DATA) +def test_resnet(batch_size, dim, res_layer_index, seq_len, kernel_size): + """Test that SEANetResnetBlock() calls are causal. Having new inputs does not change the previous output.""" + assert seq_len > kernel_size + + dilation_base = 2 + layer = SEANetResnetBlock(dim=dim, dilations=[dilation_base**res_layer_index, 1], pad_mode="constant", causal=True) + + generator = torch.Generator() + generator = generator.manual_seed(41) + layer.apply(functools.partial(_init_weights, generator=generator)) + + shape = (batch_size, dim, seq_len,) + input_hidden_states = torch.rand(shape) + + expected_output = layer(input_hidden_states) + + for end_index in range(kernel_size, seq_len + 1): + actual_output = layer(input_hidden_states[..., :end_index]) + torch.testing.assert_close(actual_output, expected_output[..., :actual_output.shape[-1]], + msg=lambda original_msg: f"Failed at end_index={end_index}: \n{original_msg}") + + +@pytest.mark.parametrize("batch_size, dim, res_layer_index, seq_len, kernel_size", SEANET_RESNET_DATA) +def test_resnet_streaming(batch_size, dim, res_layer_index, seq_len, kernel_size): + """Test that SEANetResnetBlock() streaming works as expected.""" + assert seq_len > kernel_size + + dilation_base = 2 + layer = SEANetResnetBlock(dim=dim, dilations=[dilation_base**res_layer_index, 1], pad_mode="constant", causal=True) + + generator = torch.Generator() + generator = generator.manual_seed(41) + layer.apply(functools.partial(_init_weights, generator=generator)) + + shape = (batch_size, dim, seq_len,) + input_hidden_states = torch.rand(shape) + + expected_output = layer(input_hidden_states) + + start_index = 0 + actual_outputs = [] + with layer.streaming(batch_size=batch_size): + for end_index in range(kernel_size, seq_len + 1): + actual_output = layer(input_hidden_states[..., start_index:end_index]) + start_index = end_index + actual_outputs.append(actual_output) + actual_outputs = torch.cat(actual_outputs, dim=-1) + + torch.testing.assert_close(actual_outputs, expected_output) + + +@pytest.mark.parametrize("num_timesteps", NUM_TIMESTEPS_DATA) +@pytest.mark.parametrize("seanet_kwargs", SEANET_KWARGS_DATA) +def test_nonstreaming_causal_decode(num_timesteps, seanet_kwargs): + """Test that the SEANetDecoder does not depend on future inputs.""" + + device = 'cuda' if torch.cuda.is_available() else 'cpu' + decoder = SEANetDecoder(**seanet_kwargs).to(device=device) + + generator = torch.Generator(device=device) + generator = generator.manual_seed(41) + decoder.apply(functools.partial(_init_weights, generator=generator)) + + rand_generator = torch.Generator(device=device) + rand_generator.manual_seed(2147483647) + with torch.no_grad(): + # [B, K = 8, T] + codes = torch.randn(1, seanet_kwargs['dimension'], num_timesteps, generator=rand_generator, device=device) + expected_decoded = decoder(codes) + + num_timesteps = codes.shape[-1] + for t in range(num_timesteps): + current_codes = codes[..., :t + 1] + actual_decoded = decoder(current_codes) + torch.testing.assert_close(expected_decoded[..., :actual_decoded.shape[-1]], actual_decoded, + msg=lambda original_msg: f"Failed at t={t}: \n{original_msg}") diff --git a/moshi_src/moshi/modules/streaming.py b/moshi_src/moshi/modules/streaming.py new file mode 100644 index 0000000..4193dbe --- /dev/null +++ b/moshi_src/moshi/modules/streaming.py @@ -0,0 +1,217 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +""" +Streaming module API that should be implemented by all Streaming components, +""" + +import abc +from contextlib import ExitStack +from dataclasses import dataclass +import typing as tp +from torch import nn +import torch +from ..utils.compile import CUDAGraphed + + +@dataclass +class State(abc.ABC): + """Base State for streaming, requires to be resetable and also support the context + protocol. The state __enter__ and __exit__ will be called upon entering and exiting + streaming in the parent module, but not upon `reset()` calls. + """ + + batch_size: int + device: torch.device + + def __post_init__(self): + self.exec_mask = torch.ones(self.batch_size, dtype=torch.bool, device=self.device) + self._set_exec_mask_graphed: CUDAGraphed | None = None + + def set_exec_mask(self, exec_mask: torch.Tensor): + self.exec_mask[:] = exec_mask + + def reset(self, reset_mask: torch.Tensor) -> None: + self.exec_mask[:] = torch.where(reset_mask, torch.ones_like(self.exec_mask), self.exec_mask) + + def __enter__(self) -> None: + pass + + def __exit__(self, exc_type, exc_value, traceback) -> None: + pass + + +StateT = tp.TypeVar("StateT", bound=State) + + +class StreamingModule(abc.ABC, nn.Module, tp.Generic[StateT]): + """Common API for streaming components. + + Each streaming component has a streaming state, `self._streaming_state`, which is None by default. + + To set a streaming component in streaming state, use + + with module.streaming(batch_size): + ... + + This will automatically void the streaming state when exiting the context manager. + This also automatically propagates to all streaming children module. + When the streaming state is set, modules should store whatever state they need in there. + """ + def __init__(self) -> None: + super().__init__() + self._streaming_state: StateT | None = None + self._streaming_detached: bool = False + self._cached_children: list[tuple[str, StreamingModule]] | None = None + + @property + def is_streaming(self): + return self._streaming_state is not None + + def set_streaming_detached(self, streaming_detached: bool): + """If set to False, the default, this module and all submodules will switch to streaming mode + if a parent module is set to streaming mode. + If set to True, or in detach mode, only a direct call to this module `.streaming(...)` method + will set it into streaming mode, ignoring the changes from its parents. + + This is useful if streaming over two different dimensions, e.g. for the RQ-Transformer + with the inner Depth Transformer working on the dimension of the codebooks.""" + self._streaming_detached = streaming_detached + + def _apply_named_streaming(self, fn: tp.Any): + def _handle_module(prefix: str, module: nn.Module): + if isinstance(module, StreamingModule): + # If prefix is empty, we are the direct receiver of the streaming request, + # otherwise, we are inheriting from a parent and will stop if detached. + if module._streaming_detached and prefix != "": + return + assert self._cached_children is not None + self._cached_children.append((prefix, module)) + for name, child in module.named_children(): + if prefix: + new_prefix = prefix + "." + name + else: + new_prefix = name + _handle_module(new_prefix, child) + + if self._cached_children is None: + self._cached_children = [] + _handle_module("", self) + for name, child in self._cached_children: + fn(name, child) + + def _start_streaming(self, batch_size: int, exit_stack: ExitStack): + def _start_streaming(name: str, module: StreamingModule): + assert module._streaming_state is None, f"{name} is already streaming!" + state = module._init_streaming_state(batch_size) + exit_stack.enter_context(state) + module._streaming_state = state + + self._apply_named_streaming(_start_streaming) + + def _stop_streaming(self) -> None: + def _stop_streaming(name: str, module: StreamingModule): + module._streaming_state = None + + self._apply_named_streaming(_stop_streaming) + + @abc.abstractmethod + def _init_streaming_state(self, batch_size: int) -> StateT: ... + + def streaming_forever(self, batch_size: int): + self.streaming(batch_size).__enter__() + + def streaming(self, batch_size: int) -> ExitStack: + """Context manager to enter streaming mode. Reset streaming state on exit.""" + + exit_stack = ExitStack() + self._start_streaming(batch_size, exit_stack) + exit_stack.callback(self._stop_streaming) + return exit_stack + + def reset_streaming(self, reset_mask: torch.Tensor | None = None) -> None: + """Reset the streaming state.""" + + def _reset(name: str, module: StreamingModule): + state = module._streaming_state + if state is None: + raise ValueError( + f"Trying to reset streaming, but {name} wasn't streaming." + ) + state.reset(reset_mask) + + state = self._streaming_state + assert state is not None + if reset_mask is None: + reset_mask = torch.ones(state.batch_size, device=state.device, dtype=torch.bool) + else: + reset_mask = reset_mask.to(state.device) + self._apply_named_streaming(_reset) + + def get_streaming_state(self) -> dict[str, tp.Any]: + """Return the complete streaming state, including that of sub-modules.""" + state: dict[str, tp.Any] = {} + + def _add(name: str, module: StreamingModule): + state[name] = module._streaming_state + + self._apply_named_streaming(_add) + return state + + def set_streaming_state(self, state: dict[str, tp.Any]): + """Set the streaming state, including that of sub-modules.""" + state = dict(state) + + def _set(name: str, module: StreamingModule): + if name in state: + module._streaming_state = state[name] + state.pop(name) + else: + raise RuntimeError(f"Expected to find a streaming state for {name}.") + + self._apply_named_streaming(_set) + if state: + raise RuntimeError(f"Some states were not consumed: {list(state.keys())}") + + def set_exec_mask(self, exec_mask: torch.Tensor): + """Set the execution mask, a tensor of boolean of shape `(B,), indicating + for each batch item whether the internal state should be updated or not as if + real data had been received. + + This is useful for running desynchronized streams with batching, e.g. when + the mask is False for an entry, the internal state will be unchanged by the provided + data, e.g. will be on the next step as if the previous one had never happened. + + There is no magic here, each StreamingModule subclass is responsible for respecting + the exec_mask. + """ + + state = self._streaming_state + assert state is not None + exec_mask = exec_mask.to(state.device) + + def _set_exec_mask(exec_mask: torch.Tensor): + def _set_exec_mask_fn(name: str, module: StreamingModule): + state = module._streaming_state + assert state is not None + state.set_exec_mask(exec_mask) + self._apply_named_streaming(_set_exec_mask_fn) + + if state._set_exec_mask_graphed is None: + disable = state.device.type != 'cuda' + state._set_exec_mask_graphed = CUDAGraphed(_set_exec_mask, disable=disable) + + state._set_exec_mask_graphed(exec_mask) + + +class StreamingContainer(StreamingModule[State]): + def _init_streaming_state(self, batch_size: int) -> State: + device = next(iter(self.parameters())).device + return State(batch_size, device) diff --git a/moshi_src/moshi/modules/transformer.py b/moshi_src/moshi/modules/transformer.py new file mode 100644 index 0000000..5f70ccc --- /dev/null +++ b/moshi_src/moshi/modules/transformer.py @@ -0,0 +1,956 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +""" +Transformer model, with streaming support, + CUDA Graphable. +Optimized for inference. + +See `StreamingTransformer` for more information. +""" + +from contextlib import ExitStack +from dataclasses import dataclass +import typing as tp +from einops import rearrange +import torch +import torch.nn as nn +from torch.nn import functional as F +from ..utils.compile import no_compile +from ..utils import quantize +from ..utils.quantize import replace_linear_with_qlinear +from .gating import make_gating +from .rope import RotaryEmbedding +from .streaming import StreamingModule, StreamingContainer, State +from .lora import LoRALinear +from torch.utils.checkpoint import checkpoint as torch_checkpoint + + +class LayerNormF32(nn.LayerNorm): + def forward(self, input: torch.Tensor) -> torch.Tensor: + x_f32 = input.float() + out_f32 = super().forward(x_f32) + return out_f32.to(input.dtype) + + +def _rms_norm( + x: torch.Tensor, + alpha: torch.Tensor, + dtype: tp.Optional[torch.dtype], + eps: float, +): + assert x.dim() == 3, f"RMSNorm expects 3D inputs but got {x.shape}" + x_dtype = x.dtype + if dtype is not None: + x = x.to(dtype) + var = eps + torch.mean(x**2, dim=2, keepdim=True) + y = (x * (alpha.to(var) * torch.rsqrt(var))).to(x_dtype) + return y + + +class RMSNorm(nn.Module): + def __init__( + self, + dim: int, + eps: float = 1e-5, + dtype: tp.Optional[torch.dtype] = None, + device=None, + ): + super().__init__() + self.eps = eps + self.dtype = dtype + self.alpha = nn.Parameter( + torch.full((1, 1, dim), 1.0, requires_grad=True, device=device, dtype=dtype) + ) + + def forward(self, x: torch.Tensor): + return _rms_norm(x, self.alpha, self.dtype, self.eps) + + +class LayerScale(nn.Module): + """Layer scale from [Touvron et al 2021] (https://arxiv.org/pdf/2103.17239.pdf). + This rescales diagonally the residual outputs close to 0, with a learnt scale. + + Args: + channels (int): Number of channels. + init (float): Initial scale. + channel_last (bool): If True, expect `[*, C]` shaped tensors, otherwise, `[*, C, T]`. + device (torch.device or str, optional): Device on which to initialize the module. + dtype (torch.dtype, optional): dtype to use to initialize the module. + """ + + def __init__( + self, + channels: int, + init: float = 1e-4, + channel_last: bool = True, + device=None, + dtype=None, + ): + super().__init__() + self.channel_last = channel_last + self.scale = nn.Parameter( + torch.full( + (channels,), init, requires_grad=True, device=device, dtype=dtype + ) + ) + + def forward(self, x: torch.Tensor): + if self.channel_last: + return self.scale * x + else: + return self.scale[:, None] * x + + +def create_norm_fn(norm_type: str, dim: int, **kwargs) -> nn.Module: + """Create normalization module for transformer encoder layer. + + Args: + norm_type (str): Normalization method. + dim (int): Dimension of the normalized layer. + **kwargs (dict): Additional parameters for normalization layer. + Returns: + nn.Module: Normalization module. + """ + if norm_type == "layer_norm": + return nn.LayerNorm(dim, eps=1e-5, **kwargs) + elif norm_type == "layer_norm_f32": + kwargs.pop("dtype", None) + return LayerNormF32(dim, eps=1e-8, **kwargs) + elif norm_type in {"rms_norm"}: + return RMSNorm(dim, eps=1e-5, **kwargs) + elif norm_type in {"rms_norm_f32"}: + kwargs.pop("dtype", None) + return RMSNorm(dim, eps=1e-8, dtype=torch.float, **kwargs) + else: + raise ValueError(f"Unknown norm type: {norm_type}") + + +def create_sin_embedding( + positions: torch.Tensor, + dim: int, + max_period: float = 10000, + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """Create sinusoidal positional embedding, with shape `[B, T, C]`. + + Args: + positions (torch.Tensor): LongTensor of positions. + dim (int): Dimension of the embedding. + max_period (float): Maximum period of the cosine/sine functions. + dtype (torch.dtype or str): dtype to use to generate the embedding. + Returns: + torch.Tensor: Sinusoidal positional embedding. + """ + # We aim for BTC format + assert dim % 2 == 0 + half_dim = dim // 2 + positions = positions.to(dtype) + adim = torch.arange(half_dim, device=positions.device, dtype=dtype).view(1, 1, -1) + max_period_tensor = torch.full( + [], max_period, device=positions.device, dtype=dtype + ) # avoid sync point + phase = positions / (max_period_tensor ** (adim / (half_dim - 1))) + return torch.cat([torch.cos(phase), torch.sin(phase)], dim=-1) + + +def set_attention_context(model: nn.Module, context: tp.Optional[int] = None) -> None: + """Deactivates or changes the context span (in time steps) in a model. + Args: + model (nn.Module): model over which to look for attentions. + context (int or None): new temporary context value. + + ..Note:: this is not a context manager but a plain function changing the context forever. + Initially, it was a context manager, but that led to interesting bugs when using + activation checkpointing, with the context being inconsistent between the forward + and backward. + """ + for module in model.modules(): + if isinstance(module, StreamingMultiheadAttention): + module.context = context + + +class KVCacheResult(tp.NamedTuple): + keys: torch.Tensor + values: torch.Tensor + positions: torch.Tensor + + @staticmethod + def from_kv(keys: torch.Tensor, values: torch.Tensor) -> "KVCacheResult": + B, H, T, D = keys.shape + assert tuple(values.shape[:-1]) == (B, H, T) + positions = torch.arange(T, device=keys.device, dtype=torch.long) + return KVCacheResult(keys, values, positions.expand(B, -1)) + + +class RingKVCache: + """Efficient streaming KVCache to be compatible with Cuda Graph. + + Args: + batch_size (int): Batch size. + num_heads (int): Number of heads in the attention. + dim_per_head (int): Dimension per head. + device (torch.device): Device on which to initialize the cache. + dtype (torch.dtype): dtype to use for the cache. + """ + + def __init__( + self, + batch_size: int, + num_heads: int, + dim_per_head: int, + capacity: int, + respect_exec_mask: bool = True, + device: torch.device = torch.device("cuda"), + dtype: torch.dtype = torch.bfloat16, + ): + self.capacity = capacity + self.cache = torch.zeros( + (2, batch_size, num_heads, capacity, dim_per_head), + device=device, + dtype=dtype, + ) + self.respect_exec_mask = respect_exec_mask + if self.respect_exec_mask: + self.end_offset = torch.zeros(batch_size, device=device, dtype=torch.long) + else: + self.end_offset = torch.zeros(1, device=device, dtype=torch.long) + + def reset(self, reset_mask: torch.Tensor) -> None: + self.end_offset[:] = torch.where( + reset_mask, + torch.zeros_like(self.end_offset), + self.end_offset, + ) + + def complete(self, k: torch.Tensor, v: torch.Tensor, exec_mask: torch.Tensor) -> KVCacheResult: + assert k.shape[:-1] == v.shape[:-1], (k.shape, v.shape) + B, H, T, D = k.shape + assert T > 0 + indexes = torch.arange(T, device=self.end_offset.device, dtype=self.end_offset.dtype) + indexes = indexes + self.end_offset.view(-1, 1) + indexes = indexes % self.capacity + if self.respect_exec_mask: + # indexes is [B, T] + # k is [B, H, T, D] + # cache is [B, H, T', D] + this_indexes = indexes.view(B, 1, T, 1) + this_indexes = this_indexes.expand(-1, H, T, D) + self.cache[0].scatter_(2, this_indexes, k) + self.cache[1].scatter_(2, this_indexes, v) + else: + self.cache[0].index_copy_(2, indexes[0], k) + self.cache[1].index_copy_(2, indexes[0], v) + + keys = self.cache[0] + values = self.cache[1] + + indexes = torch.arange( + self.capacity, device=self.end_offset.device, dtype=torch.long + ) + + # end_index correspond to the actual index where the last value was written. + last_offset = self.end_offset.view(-1, 1) + T - 1 + end_index = last_offset % self.capacity + delta = indexes - end_index + + # We know that if `index == end_index`, then we should output `self.end_offset`. + # If `index = end_index - 1` we should output `self.end_offset - 1` + # If `index = end_index - n` we should output `self.end_offset - n` + # Now, for `index == end_index + 1` , we actually have the oldest entry in the cache, + # so we should output `end_index + 1 - self.capacity` + + positions = torch.where( + delta <= 0, + last_offset + delta, + last_offset + delta - self.capacity, + ) + if self.respect_exec_mask: + self.end_offset[:] = torch.where( + exec_mask, + self.end_offset + T, + self.end_offset) + else: + self.end_offset.add_(T) + invalid = indexes >= self.end_offset.view(-1, 1) + positions = torch.where(invalid, torch.full_like(positions, -1), positions) + + return KVCacheResult(keys, values, positions) + + +def apply_weights_per_step(modules: nn.ModuleList, schedule: list[int] | None, + x: torch.Tensor, offset: int | None) -> torch.Tensor: + """Utility to apply a multi linear layer to the given input. A multi linear layer + applies a different set of weight for each time step. + + Args: + modules (nn.ModuleList): apply weights per step. + schedule (list[int] or None): schedule for weight sharing. + x (torch.Tensor): Input tensor, with shape `[B, T, C]`. + offset (int): offset for the current time step, in particular for decoding, with + time steps provided one by one. + """ + + if len(modules) == 1: + return modules[0](x) + + assert offset is not None, "Out of sync execution with weights per step." + + ys: list[torch.Tensor] = [] + B, T, C = x.shape + for t in range(T): + module_index = t + offset + if schedule is not None: + module_index = schedule[module_index] + y = modules[module_index](x[:, t: t + 1]) + ys.append(y) + out = torch.cat(ys, 1) + return out + + +@dataclass +class _MHAState(State): + kv_cache: RingKVCache | None + offset: torch.Tensor + offset_cpu: int + k_cross: torch.Tensor | None = None + v_cross: torch.Tensor | None = None + + def reset(self, reset_mask: torch.Tensor): + super().reset(reset_mask) + self.offset[:] = torch.where(reset_mask, torch.zeros_like(self.offset), self.offset) + if self.kv_cache is not None: + self.kv_cache.reset(reset_mask) + self.offset_cpu = 0 + + +class StreamingMultiheadAttention(StreamingModule[_MHAState]): + """Similar to `nn.MultiheadAttention` but with support for streaming, causal evaluation. + + Args: + embed_dim (int): Dimension to project to. + num_heads (int): Number of heads. + causal (bool): Causal mask applied automatically. + context (int, optional): Number of time steps the attention can access to. + When causal, can access `context` time steps into the past, and when non causal, + can access `context // 2` steps in the past, and the same in the future. + rope (`RotaryEmbedding`, optional): Rope embedding to use. + weights_per_step (int): use different weights per time step. If non zero, should correspond to the + number of possible time steps. + weights_per_step_schedule (list[int] | None): if provided, some steps will share weights when + `weights_per_step` is True, e.g. step `I` will use weights `schedule[I]`. + cross_attention (bool): True if this is to be used as a cross attention. + device (torch.device, optional): Device on which to initialize. + dtype (torch.dtype, optional): dtype to use. + """ + + _fsdp_final = True + + def __init__( + self, + embed_dim: int, + num_heads: int, + causal: bool = False, + context: tp.Optional[int] = None, + rope: tp.Optional[RotaryEmbedding] = None, + weights_per_step: int = 0, + weights_per_step_schedule: list[int] | None = None, + cross_attention: bool = False, + cache_cross_attention: bool = True, + device=None, + dtype=None, + ): + super().__init__() + factory_kwargs = {"device": device, "dtype": dtype} + + self.embed_dim = embed_dim + self.causal = causal + self.context = context + self.rope = rope + self.num_heads = num_heads + self.weights_per_step = weights_per_step + self.weights_per_step_schedule = weights_per_step_schedule + self.cross_attention = cross_attention + self.cache_cross_attention = cache_cross_attention + if cross_attention: + assert not weights_per_step, "weights_per_step not supported for cross attention." + assert rope is None, "rope and cross_attention makes no sense." + assert not causal, "causal and cross attention makes no sense." + # We do not want to activate the streaming KVCache if we are a cross attention. + # self.set_streaming_detached(True) + + out_dim = 3 * embed_dim + mult = 1 + if weights_per_step: + if weights_per_step_schedule: + assert len(weights_per_step_schedule) == weights_per_step + mult = max(weights_per_step_schedule) + 1 + else: + mult = weights_per_step + self.mult = mult + + # Split in one linear per step + self.out_projs = nn.ModuleList( + [ + nn.Linear(embed_dim, embed_dim, bias=False, **factory_kwargs) + for _ in range(mult) + ] + ) + self.in_projs = nn.ModuleList( + [ + nn.Linear(embed_dim, out_dim, bias=False, **factory_kwargs) + for _ in range(mult) + ] + ) + + self._register_load_state_dict_pre_hook(StreamingMultiheadAttention._load_hook, with_module=True) + + @staticmethod + def _load_hook(module, state_dict, prefix, *_): + mappings = { + 'in_proj_weight': 'in_projs.{i}.weight', + 'in_proj.weight': 'in_projs.{i}.weight', + 'in_proj.lora_A.weight': 'in_projs.{i}.lora_A.weight', + 'in_proj.lora_B.weight': 'in_projs.{i}.lora_B.weight', + 'out_proj.weight': 'out_projs.{i}.weight', + 'out_proj.lora_A.weight': 'out_projs.{i}.lora_A.weight', + 'out_proj.lora_B.weight': 'out_projs.{i}.lora_B.weight', + } + + mult = module.mult + # _scb suffix is for quantized data. + for suffix in ['', '_scb']: + for source, target in mappings.items(): + this_source = prefix + source + suffix + if this_source in state_dict: + weight = state_dict[this_source] + _, *OD = weight.shape + weight = weight.view(mult, -1, *OD) + for i in range(mult): + this_target = prefix + target.format(i=i) + suffix + state_dict[this_target] = weight[i] + state_dict.pop(this_source) + + def _init_streaming_state(self, batch_size: int) -> _MHAState: + in_proj = self.in_projs[0] + if isinstance(in_proj, LoRALinear): + device = in_proj.lora_A.weight.device + dtype = in_proj.lora_A.weight.dtype + elif isinstance(in_proj, nn.Linear): + device = in_proj.weight.device + dtype = in_proj.weight.dtype + elif isinstance(in_proj, quantize.QLinear): + device = in_proj.weight.device + dtype = torch.float16 + else: + raise RuntimeError(f"Unknown type {type(in_proj)} for linear.") + + dim_per_head = self.embed_dim // self.num_heads + if self.cross_attention: + kv_cache = None + else: + if self.context is None: + if self.weights_per_step: + capacity = self.weights_per_step + else: + raise RuntimeError( + "Cannot create a streaming KVCache without a context to estimate capacity." + ) + else: + capacity = self.context + + kv_cache = RingKVCache( + batch_size, self.num_heads, dim_per_head, capacity, + respect_exec_mask=not self.weights_per_step, device=device, dtype=dtype + ) + return _MHAState( + batch_size, + device, + kv_cache, + offset=torch.zeros(batch_size, device=device, dtype=torch.long), + offset_cpu=0, + ) + + def _complete_kv(self, k, v) -> KVCacheResult: + state = self._streaming_state + if state is None or state.kv_cache is None: + return KVCacheResult.from_kv(k, v) + else: + return state.kv_cache.complete(k, v, state.exec_mask) + + def _compute_cross_attention( + self, key: torch.Tensor, value: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + assert self.cross_attention + assert key is value + in_proj = self.in_projs[0] + assert in_proj.bias is None + assert isinstance(in_proj, nn.Linear) + dim = in_proj.weight.shape[0] // 3 + kv = nn.functional.linear(key, in_proj.weight[dim:]) + k, v = rearrange(kv, "b t (p h d) -> p b h t d", p=2, h=self.num_heads) + return k, v + + def update_streaming_cross_attention_src( + self, cross_attention_src: torch.Tensor) -> None: + state = self._streaming_state + assert state is not None + assert self.cross_attention + k, v = self._compute_cross_attention(cross_attention_src, cross_attention_src) + if state.k_cross is None: + state.k_cross = k + state.v_cross = v + else: + assert state.v_cross is not None + state.k_cross[:] = k + state.v_cross[:] = v + + def _get_cross_attention( + self, key: torch.Tensor, value: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + state = self._streaming_state + if state is not None and state.k_cross is not None: + assert state.v_cross is not None + return state.k_cross, state.v_cross + k, v = self._compute_cross_attention(key, value) + if state is not None and self.cache_cross_attention: + state.k_cross = k + state.v_cross = v + return k, v + + def forward(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor): + state = self._streaming_state + B, T = query.shape[:2] + + if state is None: + offset = torch.zeros(B, device=query.device, dtype=torch.long) + offset_cpu = 0 + else: + offset = state.offset + offset_cpu = state.offset_cpu + + if self.cross_attention: + assert len(self.in_projs) == 1 + in_proj = self.in_projs[0] + assert in_proj.bias is None + assert isinstance(in_proj, nn.Linear) + dim = in_proj.weight.shape[0] // 3 + q = nn.functional.linear(query, in_proj.weight[:dim]) + q = rearrange(q, "b t (h d) -> b h t d", h=self.num_heads) + k, v = self._get_cross_attention(key, value) + else: + projected = apply_weights_per_step( + self.in_projs, self.weights_per_step_schedule, query, offset_cpu) + + q, k, v = rearrange( + projected, "b t (p h d) -> p b h t d", p=3, h=self.num_heads + ) + if self.rope: + q, k = self.rope(q, k, offset, time_before_heads=False) + + k, v, pos_k = self._complete_kv(k, v) + pos_k = pos_k[:, None] + if self.causal: + pos_q = offset.view(-1, 1, 1) + torch.arange(T, device=q.device, dtype=torch.long).view( + -1, 1) + delta = pos_q - pos_k + attn_bias = (pos_k >= 0) & (delta >= 0) + if self.context is not None: + attn_bias = attn_bias & (delta < self.context) + attn_bias = attn_bias[:, None] + else: + attn_bias = None + x = F.scaled_dot_product_attention(q, k, v, attn_bias, dropout_p=0.0) + + x = rearrange(x, "b h t d -> b t (h d)") + x = apply_weights_per_step( + self.out_projs, self.weights_per_step_schedule, x, offset_cpu) + + if state is not None and not self.cross_attention: + state.offset[:] = torch.where( + state.exec_mask, + state.offset + T, + state.offset) + state.offset_cpu += T + return x + + +@dataclass +class _LayerState(State): + offset_cpu: int = 0 + + def reset(self, reset_mask: torch.Tensor): + super().reset(reset_mask) + self.offset_cpu = 0 + + +class StreamingTransformerLayer(StreamingModule[_LayerState]): + """TransformerLayer with Streaming / Causal support. + + Args: + d_model (int): Dimension of the data. + num_heads (int): Number of heads. + dim_feedforward (int): Intermediate dimension of FF module. + causal (bool): Causal mask applied automatically. + context (int, optional): Receptive field for the causal mask, infinite if None. + rope (`RotaryEmbedding`, optional): Rope embedding to use. + norm (str): Normalization to use. Currently, only 'layer_norm' is supported. + layer_scale (float, optional): If not None, LayerScale will be used with the given value as initial scale. + gating (str): if provided, replaces FFN with special gating, like GLU, GSiGLU etc. + weights_per_step (int): use different weights per time step. If non zero, should correspond to the + number of possible time steps. + weights_per_step_schedule (list[int] | None): if provided, some steps will share weights when + `weights_per_step` is True, e.g. step `I` will use weights `schedule[I]`. + skip_self_attn: If true, skips the self attention module and the norm + cross_attention (bool): If True, expect to get secondary input for cross-attention. + device (torch.device, optional): Device on which to initialize. + dtype (torch.dtype, optional): dtype to use. + """ + + _fsdp_final = True + + def __init__( + self, + d_model: int, + num_heads: int, + dim_feedforward: int | list[int] = 2048, + causal: bool = False, + context: tp.Optional[int] = None, + rope: tp.Optional[RotaryEmbedding] = None, + norm: str = "layer_norm", + layer_scale: tp.Optional[float] = None, + gating: str = "none", + weights_per_step: int = 0, + weights_per_step_schedule: list[int] | None = None, + activation=F.gelu, + skip_self_attn: bool = False, + cross_attention: bool = False, + device=None, + dtype=None, + ): + super().__init__() + factory_kwargs = {"device": device, "dtype": dtype} + # Redefine self_attn to our streaming multi-head attention + attn_kwargs: tp.Dict[str, tp.Any] = { + "embed_dim": d_model, + "num_heads": num_heads, + } + if not skip_self_attn: + self.self_attn: StreamingMultiheadAttention = StreamingMultiheadAttention( + causal=causal, + context=context, + rope=rope, + weights_per_step=weights_per_step, + weights_per_step_schedule=weights_per_step_schedule, + **attn_kwargs, # type: ignore + **factory_kwargs, # type: ignore + ) # type: ignore + self.norm1 = create_norm_fn(norm, d_model, **factory_kwargs) + self.norm2 = create_norm_fn(norm, d_model, **factory_kwargs) + # Redefine feedforward layers to expose bias parameter + self.weights_per_step = weights_per_step + self.weights_per_step_schedule = weights_per_step_schedule + self.gating: tp.Optional[nn.Module] = None + self.linear1: tp.Optional[nn.Module] = None + self.linear2: tp.Optional[nn.Module] = None + self.activation = activation + self.skip_self_attn = skip_self_attn + + num_weights = 1 + if weights_per_step is not None: + num_weights = weights_per_step + if weights_per_step_schedule is not None: + assert len(weights_per_step_schedule) == weights_per_step + num_weights = max(weights_per_step_schedule) + 1 + if isinstance(dim_feedforward, list): + assert dim_feedforward + assert len(dim_feedforward) == num_weights, ( + "Length of dim_feedforward must match weights_per_step," + f" got {len(dim_feedforward)} != {num_weights}" + ) + if gating == "none": + assert ( + not weights_per_step + ), "weights_per_step without gating not supported for now." + assert not isinstance( + dim_feedforward, list + ), "List dim_feedforward without gating not supported for now." + self.linear1 = nn.Linear( + d_model, dim_feedforward, bias=False, **factory_kwargs + ) + self.linear2 = nn.Linear( + dim_feedforward, d_model, bias=False, **factory_kwargs + ) + else: + self.linear1 = None + self.linear2 = None + if weights_per_step: + if isinstance(dim_feedforward, int): + dim_feedforward = [dim_feedforward] * num_weights + assert isinstance(dim_feedforward, list), dim_feedforward + self.gating = nn.ModuleList( + [ + make_gating(gating, d_model, dim, **factory_kwargs) + for dim in dim_feedforward + ] + ) + else: + assert isinstance(dim_feedforward, int) + self.gating = make_gating( + gating, d_model, dim_feedforward, **factory_kwargs + ) + + self.cross_attention: StreamingMultiheadAttention | None = None + if cross_attention: + self.cross_attention = StreamingMultiheadAttention( + cross_attention=True, **attn_kwargs, **factory_kwargs) # type: ignore + # Cross attention norm is always a layer norm, for no specific reason. + self.norm_cross = nn.LayerNorm(d_model, eps=1e-5, **factory_kwargs) # type: ignore + + self.layer_scale_1: nn.Module + self.layer_scale_2: nn.Module + if layer_scale is None: + self.layer_scale_1 = nn.Identity() + self.layer_scale_2 = nn.Identity() + if cross_attention: + self.layer_scale_cross = nn.Identity() + else: + self.layer_scale_1 = LayerScale(d_model, layer_scale, **factory_kwargs) # type: ignore + self.layer_scale_2 = LayerScale(d_model, layer_scale, **factory_kwargs) # type: ignore + if cross_attention: + self.layer_scale_cross = LayerScale(d_model, layer_scale, **factory_kwargs) # type: ignore + + def _init_streaming_state(self, batch_size: int) -> _LayerState: + device = next(iter(self.parameters())).device + return _LayerState(batch_size, device, offset_cpu=0) + + # feed forward block + def _ff_block(self, x: torch.Tensor) -> torch.Tensor: + state = self._streaming_state + offset = 0 + if state is not None: + offset = state.offset_cpu + x_orig = x + x = self.norm2(x) + if self.gating is None: + assert self.linear1 is not None + assert self.linear2 is not None + update = self.linear2(self.activation(self.linear1(x))) + else: + if self.weights_per_step: + assert isinstance(self.gating, nn.ModuleList) + update = apply_weights_per_step(self.gating, self.weights_per_step_schedule, x, offset) + else: + update = self.gating(x) + return x_orig.to(update) + self.layer_scale_2(update) + + def _sa_block(self, x: torch.Tensor): + if self.skip_self_attn: + return x + x_orig = x + x = self.norm1(x) + update = self.self_attn(x, x, x) + return x_orig.to(update) + self.layer_scale_1(update) + + def _cross_attention_block(self, x: torch.Tensor, + cross_attention_src: torch.Tensor) -> torch.Tensor: + assert self.cross_attention is not None + x_orig = x + x = self.norm_cross(x) + # queries are from src, keys and values from cross_attention_src. + update = self.cross_attention(x, cross_attention_src, cross_attention_src) + return x_orig + self.layer_scale_cross(update) + + def forward(self, x: torch.Tensor, cross_attention_src: torch.Tensor | None = None): + with ExitStack() as stack: + if x.device.type != 'cuda': + stack.enter_context(no_compile()) + x = self._sa_block(x) + if self.cross_attention is not None: + assert cross_attention_src is not None + x = self._cross_attention_block(x, cross_attention_src) + else: + assert cross_attention_src is None + x = self._ff_block(x) + state = self._streaming_state + if state: + state.offset_cpu += x.shape[1] + return x + + +@dataclass +class _TransformerState(State): + offsets: torch.Tensor + + def reset(self, reset_mask: torch.Tensor): + super().reset(reset_mask) + self.offsets[:] = torch.where(reset_mask, torch.zeros_like(self.offsets), self.offsets) + + +class StreamingTransformer(StreamingModule[_TransformerState]): + """Transformer with Streaming / Causal support. + + Args: + d_model (int): Dimension of the data. + num_heads (int): Number of heads. + dim_feedforward (int): Intermediate dimension of FF module. + causal (bool): Causal mask applied automatically. + context (int, optional): Receptive field for the causal mask, infinite if None. + layer_scale (float, optional): If not None, LayerScale will be used + with the given value as initial scale. + positional_embedding (str): Positional embedding strategy (sin, rope, sin_rope, or none). + max_period (float): Maximum period of the time embedding. + positional_scale (float): Scale of positional embedding, set to 0 to deactivate. + layer_class: (subclass of `StreamingTransformerLayer): class to use + to initialize the layers, allowing further customization outside of AudioCraft. + device (torch.device, optional): Device on which to initialize. + dtype (torch.dtype, optional): dtype to use. + **kwargs: See `StreamingTransformerLayer`. + """ + + def __init__( + self, + d_model: int, + num_heads: int, + num_layers: int, + dim_feedforward: int | list[int] = 2048, + causal: bool = False, + context: tp.Optional[int] = None, + positional_embedding: str = "sin", + max_period: float = 10_000, + positional_scale: float = 1.0, + betas: tp.Optional[tp.Tuple[float, float]] = None, + layer_class: tp.Type[StreamingTransformerLayer] = StreamingTransformerLayer, + quantize: bool = False, + checkpointing: bool = False, + device=None, + dtype=None, + **kwargs, + ): + super().__init__() + assert d_model % num_heads == 0 + + self.positional_embedding = positional_embedding + self.max_period = max_period + self.positional_scale = positional_scale + self.betas = betas + + assert positional_embedding in {"sin", "rope", "sin_rope", "none"} + self.rope: tp.Optional[RotaryEmbedding] = None + if self.positional_embedding in {"rope", "sin_rope"}: + self.rope = RotaryEmbedding(max_period=max_period) + + self.checkpointing = checkpointing + + self.layers = nn.ModuleList() + for _ in range(num_layers): + self.layers.append( + layer_class( + d_model=d_model, + num_heads=num_heads, + dim_feedforward=dim_feedforward, + causal=causal, + context=context, + rope=self.rope, + device=device, + dtype=dtype, + **kwargs, + ) + ) + if quantize: + # Quantizing layers one by one to avoid taking too much space during init. + self.layers[-1].to(device=device, dtype=dtype) + replace_linear_with_qlinear(self.layers[-1]) + + def _init_streaming_state(self, batch_size: int) -> _TransformerState: + device = next(self.parameters()).device + return _TransformerState(batch_size, device, offsets=torch.zeros(batch_size, device=device, dtype=torch.long)) + + def forward(self, x: torch.Tensor, *args, **kwargs): + B, T, C = x.shape + + dtype_input = x.dtype + state = self._streaming_state + if state is None: + offsets = torch.zeros(1, dtype=torch.long, device=x.device) + else: + offsets = state.offsets + + if self.positional_embedding in {"sin", "sin_rope"}: + positions = torch.arange(T, device=x.device).view(1, -1, 1) + positions = positions + offsets.view(-1, 1, 1) + pos_emb = create_sin_embedding( + positions, C, max_period=self.max_period, dtype=x.dtype + ) + x = x + self.positional_scale * pos_emb + + for layer in self.layers: + if self.checkpointing: + y = torch_checkpoint( + layer, x, *args, use_reentrant=False, + determinism_check='none', + preserve_rng_state=False, + **kwargs) + assert isinstance(y, torch.Tensor) + x = y + else: + x = layer(x, *args, **kwargs) + + if state is not None: + state.offsets[:] = torch.where( + state.exec_mask, + state.offsets + T, + state.offsets) + return x.to(dtype_input) + + +class ProjectedTransformer(StreamingContainer): + """Transformer with optional projections of the input and output to different dimensions when needed. + Supports multiple outputs. + + Args: + input_dimension (int): dimension of the input. + output_dimensions (tuple[int]): dimensions of the outputs. + d_model (int): inner dimension of the Transformer. + conv_layout (bool): If True, expects `[B, C, T]` shaped tensors, otherwise, `[B, T, C]`. + Similarly, the output will have the same layout. + """ + + def __init__( + self, + input_dimension: int, + output_dimensions: tp.Tuple[int, ...], + d_model: int, + *, + conv_layout: bool = False, + **kwargs, + ): + super().__init__() + self.transformer = StreamingTransformer(d_model=d_model, **kwargs) + self.input_dimension = input_dimension + self.output_dimensions = output_dimensions + self.conv_layout = conv_layout + self.input_proj = None + if d_model != input_dimension: + self.input_proj = nn.Linear(input_dimension, d_model, bias=False) + + self.output_projs = nn.ModuleList() + for output_dimension in output_dimensions: + if d_model == output_dimension: + self.output_projs.append(nn.Identity()) + else: + self.output_projs.append( + nn.Linear(d_model, output_dimension, bias=False) + ) + + def forward(self, x, *args, **kwargs): + if self.conv_layout: + x = x.transpose(1, 2) + if self.input_proj is not None: + x = self.input_proj(x) + z = self.transformer(x, *args, **kwargs) + ys = [] + for output_proj in self.output_projs: + y = output_proj(z) + if self.conv_layout: + y = y.transpose(1, 2) + ys.append(y) + return ys diff --git a/moshi_src/moshi/quantization/__init__.py b/moshi_src/moshi/quantization/__init__.py new file mode 100644 index 0000000..0e5e1b3 --- /dev/null +++ b/moshi_src/moshi/quantization/__init__.py @@ -0,0 +1,13 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +"""RVQ.""" +# flake8: noqa +from .vq import ResidualVectorQuantizer, SplitResidualVectorQuantizer +from .base import BaseQuantizer, DummyQuantizer, QuantizedResult diff --git a/moshi_src/moshi/quantization/base.py b/moshi_src/moshi/quantization/base.py new file mode 100644 index 0000000..242d5f9 --- /dev/null +++ b/moshi_src/moshi/quantization/base.py @@ -0,0 +1,170 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +""" +Base class for all quantizers. +""" + +from dataclasses import dataclass, field +import typing as tp + +import torch +from torch import nn + + +@dataclass +class QuantizedResult: + x: torch.Tensor + codes: torch.Tensor + bandwidth: torch.Tensor # bandwidth in kb/s used, per batch item. + penalty: tp.Optional[torch.Tensor] = None + metrics: dict = field(default_factory=dict) + + +class BaseQuantizer(nn.Module): + """Base class for quantizers.""" + + def __init__(self): + super().__init__() + self._ema_frozen = False + + def forward(self, x: torch.Tensor, frame_rate: int) -> QuantizedResult: + """ + Given input tensor x, returns first the quantized (or approximately quantized) + representation along with quantized codes, bandwidth, and any penalty term for the loss. + Finally, this returns a dict of metrics to update logging etc. + Frame rate must be passed so that the bandwidth is properly computed. + """ + raise NotImplementedError() + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """Encode a given input tensor with the specified sample rate at the given bandwidth.""" + raise NotImplementedError() + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """Decode the given codes to the quantized representation.""" + raise NotImplementedError() + + @property + def cardinality(self) -> int: + """Cardinality of each codebook.""" + raise NotImplementedError() + + @property + def total_codebooks(self) -> int: + """Total number of codebooks.""" + raise NotImplementedError() + + @property + def num_codebooks(self) -> int: + """Number of active codebooks.""" + raise NotImplementedError() + + @property + def semantic_quantizer(self) -> 'BaseQuantizer': + """This returns the quantizer that models the first level of the hierarchy (typically semantic). + + In this case, it's the quantizer itself. + """ + return self + + @property + def acoustic_quantizer(self) -> 'BaseQuantizer': + """This returns the quantizer that models the higher levels of the hierarchy (typically acoustic). + + In this case, it's the quantizer itself. + """ + return self + + def set_num_codebooks(self, n: int) -> None: + """Set the number of active codebooks.""" + raise NotImplementedError() + + @property + def ema_frozen(self) -> bool: + """Whether to apply ema to the codebooks.""" + return self._ema_frozen + + def ema_frozen_(self, ema_frozen: bool) -> None: + """Set whether ema should be applied to the codebooks.""" + self._ema_frozen = ema_frozen + + +class DummyQuantizer(BaseQuantizer): + """Fake quantizer that actually does not perform any quantization.""" + + def __init__( + self, + dimension: int, + input_dimension: tp.Optional[int] = None, + output_dimension: tp.Optional[int] = None, + ): + super().__init__() + self.dimension = dimension + self.input_dimension = input_dimension or dimension + self.output_dimension = output_dimension or dimension + self.input_proj: torch.nn.Module + self.output_proj: torch.nn.Module + if self.input_dimension == self.dimension: + self.input_proj = torch.nn.Identity() + else: + self.input_proj = torch.nn.Conv1d( + self.input_dimension, self.dimension, 1, bias=False + ) + if self.input_dimension == self.dimension: + self.output_proj = torch.nn.Identity() + else: + self.output_proj = torch.nn.Conv1d( + self.dimension, self.output_dimension, 1, bias=False + ) + + def forward(self, x: torch.Tensor, frame_rate: int): + q = x.unsqueeze(1) + x = self.output_proj(self.input_proj(x)) + return QuantizedResult( + x, q, torch.tensor(q.numel() * 32 * frame_rate / 1000 / len(x)).to(x) + ) + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """Encode a given input tensor with the specified sample rate at the given bandwidth. + In the case of the DummyQuantizer, the codes are actually identical + to the input and resulting quantized representation as no quantization is done. + """ + x = self.input_proj(x) + return x.unsqueeze(1) + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """Decode the given codes to the quantized representation. + In the case of the DummyQuantizer, the codes are actually identical + to the input and resulting quantized representation as no quantization is done. + """ + y = codes.squeeze(1) + return self.output_proj(y) + + @property + def total_codebooks(self): + """Total number of codebooks.""" + return 1 + + @property + def num_codebooks(self): + """Total number of codebooks.""" + return self.total_codebooks + + def set_num_codebooks(self, n: int): + """Set the number of active codebooks.""" + raise AttributeError( + "Cannot override the number of codebooks for the dummy quantizer" + ) + + @property + def cardinality(self) -> int: + """Cardinality of each codebook.""" + return 1 diff --git a/moshi_src/moshi/quantization/core_vq.py b/moshi_src/moshi/quantization/core_vq.py new file mode 100644 index 0000000..09f9bbd --- /dev/null +++ b/moshi_src/moshi/quantization/core_vq.py @@ -0,0 +1,528 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import math +import typing as tp + +from einops import rearrange, repeat +import torch +from torch import nn +from torch import distributed +import torch.nn.functional as F + + +class _CodebookForwardResult(tp.NamedTuple): + quantized: torch.Tensor + codes: torch.Tensor + metrics: tp.Dict[str, torch.Tensor] + + +class _VQForwardResult(tp.NamedTuple): + quantized: torch.Tensor + codes: torch.Tensor + loss: torch.Tensor + metrics: tp.Dict[str, torch.Tensor] + + +def _ema_inplace(moving_avg: torch.Tensor, new: torch.Tensor, decay: float) -> None: + moving_avg.data.mul_(decay).add_(new, alpha=(1 - decay)) + + +def _sample_vectors(samples: torch.Tensor, num: int) -> torch.Tensor: + num_samples, device = samples.shape[0], samples.device + + if num_samples >= num: + indices = torch.randperm(num_samples, device=device)[:num] + else: + indices = torch.randint(0, num_samples, (num,), device=device) + + return samples[indices] + + +def _compute_entropy(usage: torch.Tensor) -> torch.Tensor: + # Usage is some unnormalized distribution. + proba = usage / usage.sum() + p_log_p = torch.where( + proba == 0, zero_scalar(usage.device), proba * torch.log(proba) + ) + return -p_log_p.sum() + + +def _is_distributed() -> bool: + # Checks if we need to use distributed routines. + return distributed.is_initialized() and distributed.get_world_size() > 1 + + +def _average_tensors(tensors: tp.Sequence[torch.Tensor]) -> None: + if not _is_distributed(): + return + world_size = distributed.get_world_size() + handles = [] + for tensor in tensors: + handle = distributed.all_reduce( + tensor.data, op=distributed.ReduceOp.SUM, async_op=True) + handles.append(handle) + for tensor, handle in zip(tensors, handles): + handle.wait() + tensor.data /= world_size + + +def _run_kmeans(samples: torch.Tensor, num_clusters: int, num_iters: int = 50) -> tp.Tuple[torch.Tensor, torch.Tensor]: + # Kmeans algorithm used to initialize the codebooks. + dim = samples.shape[-1] + means = _sample_vectors(samples, num_clusters) + bins = None + + for _ in range(num_iters): + dists = torch.cdist(samples[None], means[None], p=2)[0] + buckets = dists.argmin(dim=-1) + bins = torch.bincount(buckets, minlength=num_clusters) + zero_mask = bins == 0 + bins.clamp_(min=1) + + new_means = torch.zeros_like(means) + new_means.scatter_add_(0, repeat(buckets, "n -> n d", d=dim), samples) + new_means /= bins[..., None] + resampled = _sample_vectors(samples, num_clusters) + means = torch.where(zero_mask[..., None], resampled, new_means) + + assert bins is not None + return means, bins + + +def zero_scalar(device) -> torch.Tensor: + """Returns a 0. value on the given device without introducing a synchronization point.""" + return torch.zeros([1], device=device)[0] + + +class EuclideanCodebook(nn.Module): + """Codebook with Euclidean distance. + + Args: + dim (int): Dimension. + codebook_size (int): Codebook size. + decay (float): Decay for exponential moving average over the codebooks. + epsilon (float): Epsilon value for numerical stability. + threshold_usage_ratio (float): Defines the threshold for the cluster usage under which a centroid + is replaced. This is expressed as a fraction of the usage a centroid would get under + a uniform distribution, so that it doesn't depend on the batch size etc. + replaced_usage_ratio (float): When replacing a centroid, use this as an initial centroid usage, + to avoid the centroid getting replaced too quickly. + check_unused_every (int): Check for unused centroids every `check_unused_every` iterations. + This is to avoid too many synchronization points. + + Buffers: + cluster_usage (torch.Tensor): EMA of the cluster usage per batch, e.g. this will + be dependent on the batch size etc. + embedding_sum (torch.Tensor): EMA of the sum of the assigned points to each cluster. + In particular, this can be normalized by `cluster_usage` to obtain the + actual cluster centroids. + """ + + def __init__( + self, + dim: int, + codebook_size: int, + decay: float = 0.99, + epsilon: float = 1e-5, + threshold_usage_ratio: float = 0.1, + replaced_usage_ratio: float = 1.0, + check_unused_every: int = 5, + ): + super().__init__() + self.decay = decay + + self.dim = dim + self.codebook_size = codebook_size + + self.epsilon = epsilon + self.threshold_usage_ratio = threshold_usage_ratio + self.replaced_usage_ratio = replaced_usage_ratio + self.check_unused_every = check_unused_every + self._next_unused_check = check_unused_every + self._cached_initialized = False + + self._initialized: torch.Tensor + self.cluster_usage: torch.Tensor + self.embedding_sum: torch.Tensor + self._embedding: torch.Tensor + self.register_buffer("_initialized", torch.tensor([False], dtype=torch.float)) + self.register_buffer("cluster_usage", torch.ones(codebook_size)) + embedding = torch.zeros(codebook_size, dim) + self.register_buffer("embedding_sum", embedding) + self.register_buffer("_embedding", None, persistent=False) + + def _load_from_state_dict(self, state_dict, prefix, *args, **kwargs) -> None: + # Mapping old names to new names + mappings = { + "inited": "_initialized", + "cluster_size": "cluster_usage", + "embed_avg": "embedding_sum", + "embed_sum": "embedding_sum", + } + for old_name, new_name in mappings.items(): + old_name = prefix + old_name + if old_name in state_dict: + value = state_dict.pop(old_name) + if new_name is not None: + state_dict[prefix + new_name] = value + super()._load_from_state_dict(state_dict, prefix, *args, **kwargs) + + @property + def embedding(self) -> torch.Tensor: + if self._embedding is None: + embedding = ( + self.embedding_sum / self.cluster_usage.clamp(min=self.epsilon)[:, None] + ) + self.register_buffer("_embedding", embedding, persistent=False) + return embedding + return self._embedding + + @property + def initialized(self) -> bool: + """Cached version of self._initialized, + This assumes that once the module is initialized, it will never go back to the uninitialized state.""" + if not self._cached_initialized: + self._cached_initialized = bool(self._initialized.item()) + return self._cached_initialized + + def _init_embedding(self, data: torch.Tensor) -> None: + # Initialize the codebook, e.g. using kmeans. + if self.initialized: + return + + rank = 0 + if _is_distributed(): + rank = distributed.get_rank() + # First gathering shapes in case not all GPUs have the same effective batch size. + # then gathering the actual content. + if rank == 0: + other_shapes: tp.List[torch.Size] = [None] * distributed.get_world_size() # type: ignore + distributed.gather_object(data.shape, other_shapes) + other_data: tp.List[torch.Tensor] = [ + torch.empty(shape, device=data.device, dtype=data.dtype) for shape in other_shapes] + distributed.gather(data, other_data) + data = torch.cat(other_data, dim=0) + else: + distributed.gather_object(data.shape) + distributed.gather(data) + if rank == 0: + embedding, cluster_usage = _run_kmeans(data, self.codebook_size) + self.embedding_sum.data.copy_(embedding * cluster_usage[:, None]) + self.cluster_usage.data.copy_(cluster_usage) + self._initialized.data.fill_(1) + # Make sure all buffers across workers are in sync after initialization + self._broadcast_buffers() + + def _broadcast_buffers(self) -> None: + if _is_distributed(): + for buffer in self.buffers(): + distributed.broadcast(buffer, 0) + + def _replace_expired_codes(self, samples: torch.Tensor, mask: torch.Tensor) -> None: + # Replaces expired centroids, as indicated by `mask` (a true value indicate the code needs to be replaced). + # The new codes are sampled from the batch `samples`. + new_vectors = _sample_vectors(samples, self.codebook_size) + replace_cluster_usage = ( + self.replaced_usage_ratio * self.cluster_usage.sum() / self.codebook_size + ) + self.embedding_sum[:] = torch.where( + mask[:, None], replace_cluster_usage * new_vectors, self.embedding_sum + ) + self.cluster_usage[:] = torch.where( + mask, replace_cluster_usage, self.cluster_usage + ) + + def _check_expired_codes(self, batch_samples: torch.Tensor) -> torch.Tensor: + # Checks whether some centroids are under utilized, and replace them if necessary. + if not self.initialized: + return zero_scalar(batch_samples.device) + + self._next_unused_check -= 1 + if self._next_unused_check > 0: + return zero_scalar(batch_samples.device) + # we don't check every iteration to avoid having too many sync points. + self._next_unused_check = self.check_unused_every + threshold_cluster_usage = self.threshold_usage_ratio * self.cluster_usage.sum() / self.codebook_size + expired_codes = self.cluster_usage < threshold_cluster_usage + + assert batch_samples.dim() == 2 + self._replace_expired_codes(batch_samples, mask=expired_codes) + self._broadcast_buffers() + + return expired_codes.float().mean() + + def _reshape_input(self, x: torch.Tensor) -> torch.Tensor: + # Flattens all the dimensions but the last one, e.g. return a vector of shape `[N, D]`. + x = rearrange(x, "... d -> (...) d") + return x + + def _reshape_codes(self, codes: torch.Tensor, shape: torch.Size) -> torch.Tensor: + return codes.view(*shape[:-1]) + + def _quantize(self, x: torch.Tensor) -> torch.Tensor: + # Projects each vector in `x` over the nearest centroid and return its index. + # `x` should be `[N, D]` with `N` the number of input vectors and `D` the dimension. + assert x.dim() == 2 + dists = torch.cdist(x[None], self.embedding[None], p=2)[0] + codes = dists.argmin(dim=-1) + return codes + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """Given a tensor `x` of shape `[*, D]`, returns a tensor of integer codes of shape `[*]`. + The codes are defined as the indexes of the centroids nearest to each vector in `x`. + """ + assert x.dtype.is_floating_point, f"Input should be floats, got {x.dtype}" + shape = x.shape + x = self._reshape_input(x) + codes = self._quantize(x) + codes = self._reshape_codes(codes, shape) + return codes + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """Given a tensor of codes of shape `[*]`, returns a tensor of shape `[*, D]`, + corresponding to the centroids associated to each code index. + """ + assert ( + not codes.dtype.is_floating_point + ), f"Codes should be integers, got {codes.dtype}" + quantized = F.embedding(codes, self.embedding) + return quantized + + def forward( + self, x: torch.Tensor, initialize: bool = True + ) -> _CodebookForwardResult: + shape = x.shape + x = self._reshape_input(x) + + if self.training and initialize: + # If initialize is False, we are not allowed to initialize this layer + # and the rest of the code will operate on a 0 filled codebook. + # This is due to previous layers having used the batch to run kmeans init + # and thus, the residuals are mostly 0s. + self._init_embedding(x.detach()) + + flat_codes = self._quantize(x) + codes = self._reshape_codes(flat_codes, shape) + quantized = self.decode(codes) + metrics: tp.Dict[str, torch.Tensor] = {} + + if self.training: + # We do the expiry of the unused codes at this point as buffers are in sync + # and all the workers will take the same decision. + expired = self._check_expired_codes(x) + metrics['rvq_expired'] = expired + cluster_usage = torch.zeros_like(self.cluster_usage) + cluster_usage.scatter_add_( + 0, flat_codes, torch.ones_like(flat_codes, dtype=cluster_usage.dtype)) + _ema_inplace(self.cluster_usage, cluster_usage, self.decay) + + if self.initialized: + # We report the entropy normalized by that of the uniform distribution, + # This means the codebooks are optimally used when entropy=1. + metrics['rvq_entropy'] = _compute_entropy(self.cluster_usage) / math.log(self.codebook_size) + + embedding_sum = torch.zeros_like(self.embedding_sum) + embedding_sum.scatter_add_(0, repeat(flat_codes, "n -> n d", d=self.dim), x) + _ema_inplace(self.embedding_sum, embedding_sum, self.decay) + self.register_buffer('_embedding', None) + + return _CodebookForwardResult(quantized, codes, metrics) + + +class VectorQuantization(nn.Module): + """Vector quantization implementation. + Currently supports only euclidean distance. + + Args: + dim (int): Dimension + codebook_size (int): Codebook size + codebook_dim (int): Codebook dimension. If not defined, uses the specified dimension in dim. + decay (float): Decay for exponential moving average over the codebooks. + epsilon (float): Epsilon value for numerical stability. + threshold_usage_ratio (float): Defines the threshold for the cluster usage under which a centroid + is replaced. This is expressed as a fraction of the usage a centroid would get under + a uniform distribution, so that it doesn't depend on the batch size etc. + replaced_usage_ratio (float): When replacing a centroid, use this as an initial centroid usage, + to avoid the centroid getting replaced too quickly. + check_unused_every (int): Check for unused centroids every `check_unused_every` iterations. + This is to avoid too many synchronization points. + """ + + def __init__( + self, + dim: int, + codebook_size: int, + codebook_dim: tp.Optional[int] = None, + decay: float = 0.99, + epsilon: float = 1e-5, + threshold_usage_ratio: float = 0.1, + **kwargs, + ): + super().__init__() + if codebook_dim is None: + codebook_dim = dim + + requires_projection = codebook_dim != dim + self.project_in = ( + nn.Linear(dim, codebook_dim) if requires_projection else nn.Identity() + ) + self.project_out = ( + nn.Linear(codebook_dim, dim) if requires_projection else nn.Identity() + ) + self.epsilon = epsilon + self._codebook = EuclideanCodebook( + dim=codebook_dim, + codebook_size=codebook_size, + decay=decay, + epsilon=epsilon, + threshold_usage_ratio=threshold_usage_ratio, + **kwargs, + ) + self.codebook_size = codebook_size + + @property + def embedding(self): + return self._codebook.embedding + + @property + def initialized(self): + return self._codebook.initialized + + def _rearrange_input(self, x): + x = rearrange(x, "b d n -> b n d") + return x + + def _rearrange_output(self, quantized): + quantized = rearrange(quantized, "b n d -> b d n") + return quantized + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """Encodes `x` into discrete integer codes.""" + x = self._rearrange_input(x) + x = self.project_in(x) + codes = self._codebook.encode(x) + return codes + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """Converts integer codes into quantized vectors.""" + quantized = self._codebook.decode(codes) + quantized = self.project_out(quantized) + quantized = self._rearrange_output(quantized) + return quantized + + def forward(self, x: torch.Tensor, initialize: bool = True) -> _VQForwardResult: + x = self._rearrange_input(x) + quantized, codes, metrics = self._codebook(x, initialize=initialize) + + if self.training: + quantized = x + (quantized - x).detach() + loss = F.mse_loss(x, quantized.detach()) + else: + loss = zero_scalar(x.device) + + quantized = self.project_out(quantized) + quantized = self._rearrange_output(quantized) + + return _VQForwardResult(quantized, codes, loss, metrics) + + +class ResidualVectorQuantization(nn.Module): + """Residual vector quantization implementation. + + Follows Algorithm 1. in https://arxiv.org/pdf/2107.03312.pdf + """ + + def __init__(self, *, num_quantizers: int, codebook_offset: int, **kwargs): + super().__init__() + self.layers = nn.ModuleList( + [VectorQuantization(**kwargs) for _ in range(num_quantizers)] + ) + self.codebook_offset = codebook_offset + + def forward( + self, x: torch.Tensor, n_q: tp.Optional[int] = None + ) -> _VQForwardResult: + """ + Args: + x (torch.Tensor): input tensor to quantize, of shape `[B, C, T]`. + n_q (int or None): if provided, number of codebook levels to use in RVQ. + """ + + quantized_out = zero_scalar(x.device) + residual = x + + all_losses = [] + all_codes = [] + all_metrics: tp.Dict[str, torch.Tensor] = {} + + n_q = n_q or len(self.layers) + previous_layer_is_initialized = True + + for i, layer in enumerate(self.layers[:n_q]): # type: ignore + if self.training: + this_layer_is_initialized = layer.initialized + # We only allow the kmeans initialization if the previous layer is already initialized from the previous + # iterations, this is to avoid learning the subsequent kmeans on the same batch, which would eventually + # lead to its exhaustion and running kmeans on 0 values. + quantized, codes, loss, metrics = layer( + residual, initialize=previous_layer_is_initialized + ) + if self.training: + previous_layer_is_initialized = this_layer_is_initialized # type: ignore + + quantized = quantized.detach() + residual = residual - quantized + quantized_out = quantized_out + quantized + + all_codes.append(codes) + all_losses.append(loss) + + for key, value in metrics.items(): + if key in all_metrics: + all_metrics[key] += value / n_q + else: + all_metrics[key] = value / n_q + all_metrics[key + f"_{i + self.codebook_offset}"] = value + + if self.training: + # Solving subtle bug with STE and RVQ: https://github.com/facebookresearch/encodec/issues/25 + quantized_out = x + (quantized_out - x).detach() + to_average = [] + for layer in self.layers: + assert isinstance(layer, VectorQuantization) + to_average += [layer._codebook.cluster_usage, layer._codebook.embedding_sum] + _average_tensors(to_average) + + out_losses, out_codes = map(torch.stack, (all_losses, all_codes)) + return _VQForwardResult(quantized_out, out_codes, out_losses, all_metrics) + + def encode(self, x: torch.Tensor, n_q: tp.Optional[int] = None) -> torch.Tensor: + """Encodes `x` into discrete integer codes. If `n_q` is provided, only uses the first `n_q` codebook levels.""" + residual = x + all_indices = [] + n_q = n_q or len(self.layers) + for layer in self.layers[:n_q]: # type: ignore + assert isinstance(layer, VectorQuantization) + indices = layer.encode(residual) + quantized = layer.decode(indices) + residual = residual - quantized + all_indices.append(indices) + out_indices = torch.stack(all_indices) + return out_indices + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """Converts the integer codes into quantized vectors.""" + quantized = zero_scalar(codes.device) + for idx, layer_codes in enumerate(codes): + layer = self.layers[idx] + assert isinstance(layer, VectorQuantization) + quantized = quantized + layer.decode(layer_codes) + return quantized diff --git a/moshi_src/moshi/quantization/vq.py b/moshi_src/moshi/quantization/vq.py new file mode 100644 index 0000000..bb39727 --- /dev/null +++ b/moshi_src/moshi/quantization/vq.py @@ -0,0 +1,318 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import math +import random +import typing as tp + +import torch + +from .base import BaseQuantizer, QuantizedResult +from .core_vq import ResidualVectorQuantization + + +class ResidualVectorQuantizer(BaseQuantizer): + """Residual Vector Quantizer. + + Args: + dimension (int): Dimension of the codebooks. + input_dimension (None or int): dimension of the input, defaults to `dimension` if not provided. + output_dimension (None or int): dimension of the output, defaults to `dimension` if not provided. + n_q (int): Number of vector quantizers used. + q_dropout (bool): Random quantizer drop out at train time. + no_quantization_rate (float): Gives the probability of applying no quantization at all + at train time. The RVQ codebooks will still get the input value to learn the proper codebook. + bins (int): Codebook size. + decay (float): Decay for exponential moving average over the codebooks. + threshold_usage_ratio (float): Defines the threshold for the cluster usage under which a centroid + is replaced. This is expressed as a fraction of the usage a centroid would get under + a uniform distribution, so that it doesn't depend on the batch size etc. + replaced_usage_ratio (float): When replacing a centroid, use this as an initial centroid usage, + to avoid the centroid getting replaced too quickly. + codebook_offset (int): Offset to use for the codebook indices. This is useful when using multiple quantizers + such as in SplitResidualVectorQuantizer. + force_projection (bool): Whether to force input and output projections even when dimension is constant. + generator_seed (int or None): seed used to initialize the RNG used for no quantization. + """ + + def __init__( + self, + dimension: int = 128, + input_dimension: tp.Optional[int] = None, + output_dimension: tp.Optional[int] = None, + n_q: int = 8, + q_dropout: bool = False, + no_quantization_rate: float = 0.0, + bins: int = 1024, + decay: float = 0.99, + threshold_usage_ratio: float = 0.1, + replaced_usage_ratio: float = 1.0, + codebook_offset: int = 0, + force_projection: bool = False, + ): + super().__init__() + self.max_n_q = n_q + self.n_q = n_q + self.q_dropout = q_dropout + self.no_quantization_rate = no_quantization_rate + self.dimension = dimension + self.input_dimension = input_dimension or dimension + self.output_dimension = output_dimension or dimension + self.bins = bins + self.decay = decay + self.rng_dropout = random.Random(1234) + self.input_proj: torch.nn.Module + self.output_proj: torch.nn.Module + if self.input_dimension == self.dimension and not force_projection: + self.input_proj = torch.nn.Identity() + else: + self.input_proj = torch.nn.Conv1d( + self.input_dimension, self.dimension, 1, bias=False + ) + if self.output_dimension == self.dimension and not force_projection: + self.output_proj = torch.nn.Identity() + else: + self.output_proj = torch.nn.Conv1d( + self.dimension, self.output_dimension, 1, bias=False + ) + self.vq = ResidualVectorQuantization( + dim=self.dimension, + codebook_size=self.bins, + num_quantizers=self.n_q, + decay=self.decay, + threshold_usage_ratio=threshold_usage_ratio, + replaced_usage_ratio=replaced_usage_ratio, + codebook_offset=codebook_offset, + ) + + def forward(self, x: torch.Tensor, frame_rate: int): + """ + Args: + x (torch.Tensor): Input tensor of shape [B, C, T] with `C` number of channels. + frame_rate (int): frame rate of the input (e.g `T = frame_rate * duration`), used to compute + the bandwidth. + + Returns: + QuantizedResult: Quantized result with the following attributes: + - `x` (torch.Tensor): Quantized tensor of shape [B, C, T]. + - `codes` (torch.Tensor): Quantized codes of shape [B, K, T] with `K` number of codebooks. + - `bw` (torch.Tensor): Bandwidth of the quantized tensor in kbits per second. + - `penalty` (torch.Tensor): Commitment loss. + - `metrics` (dict): RVQ metrics, in particular rate of dead code replacement, and entropy. + """ + n_q = self.n_q + x = self.input_proj(x) + if self.training and self.q_dropout: + n_q = self.rng_dropout.randint(1, self.n_q) + bw_per_q = math.log2(self.bins) * frame_rate / 1000 + quantized, codes, commit_loss, metrics = self.vq(x, n_q=n_q) + B, _, _ = quantized.shape + if self.training and self.no_quantization_rate > 0: + mask = (torch.rand(B, 1, 1, device=x.device) <= self.no_quantization_rate).float() + quantized = x * mask + (1 - mask) * quantized + quantized = self.output_proj(quantized) + codes = codes.transpose(0, 1) + # codes is [B, K, T], with T frames, K nb of codebooks. + bw = torch.tensor(n_q * bw_per_q).to(x) + return QuantizedResult(quantized, codes, bw, penalty=torch.mean(commit_loss), metrics=metrics) + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """Encode a given input tensor with the specified frame rate at the given bandwidth. + The RVQ encode method sets the appropriate number of quantizer to use + and returns indices for each quantizer. + """ + n_q = self.n_q + if x.shape[-1] == 0: + return torch.empty((x.shape[0], n_q, 0), device=x.device, dtype=torch.int64) + + x = self.input_proj(x) + codes = self.vq.encode(x, n_q=n_q) + codes = codes.transpose(0, 1) + # codes is [B, K, T], with T frames, K nb of codebooks. + return codes + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """Decode the given codes to the quantized representation.""" + # codes is [B, K, T], with T frames, K nb of codebooks, vq.decode expects [K, B, T]. + codes = codes.transpose(0, 1) + quantized = self.vq.decode(codes) + quantized = self.output_proj(quantized) + return quantized + + @property + def total_codebooks(self): + return self.max_n_q + + @property + def num_codebooks(self): + return self.n_q + + def set_num_codebooks(self, n: int): + assert n >= 0 and n <= self.max_n_q + self.n_q = n + + @property + def cardinality(self) -> int: + return self.bins + + +class SplitResidualVectorQuantizer(BaseQuantizer): + """Residual Vector Quantizer with separate projections for the first quantizer and the rest. + + Args: + n_q (int): Number of residual vector quantizers used. + n_semantic_q (int): Number of residual vector quantizers used for the semantic quantizer. + **kwargs: Arguments to the constructor of `ResidualVectorQuantizer` that are shared between both. + """ + + def __init__( + self, + *, + n_q: int = 8, + n_q_semantic: int = 1, + **kwargs, + ): + super().__init__() + assert n_q > n_q_semantic, ( + f"Number of quantizers {n_q} must be larger " + f"than the number of semantic quantizers {n_q_semantic}." + ) + self.max_n_q = n_q + self.n_q_semantic = n_q_semantic + self.n_q_acoustic = n_q - n_q_semantic + q_dropout = kwargs.pop("q_dropout", False) + self.rvq_first = ResidualVectorQuantizer( + n_q=n_q_semantic, force_projection=True, q_dropout=False, **kwargs + ) + self.rvq_rest = ResidualVectorQuantizer( + n_q=n_q - n_q_semantic, + codebook_offset=1, + force_projection=True, + q_dropout=q_dropout, + **kwargs, + ) + + def _renorm_and_add( + self, + first_val: torch.Tensor, + rest_val: torch.Tensor, + n_q_semantic: int, + n_q_acoustic: int, + ): + """Renormalizes values from `rvq_first` and `rvq_rest` and adds them. + + This allows correcting statistics that are normalized by the number of quantizers. To renormalize, we use the + number of quantizers that are actually used, e.g. taking into account quantizer dropout. + """ + n_q = n_q_semantic + n_q_acoustic + renorm_first_val = first_val * n_q_semantic / n_q + renorm_rest_val = rest_val * n_q_acoustic / n_q + return renorm_first_val + renorm_rest_val + + def forward(self, x: torch.Tensor, frame_rate: int): + """ + Args: + x (torch.Tensor): Input tensor of shape [B, C, T] with `C` number of channels. + frame_rate (int): frame rate of the input (e.g `T = frame_rate * duration`), used to compute + the bandwidth. + + Returns: + QuantizedResult: Quantized result with the following attributes: + - `x` (torch.Tensor): Quantized tensor of shape [B, C, T]. + - `codes` (torch.Tensor): Quantized codes of shape [B, K, T] with `K` number of codebooks. + - `bw` (torch.Tensor): Bandwidth of the quantized tensor in kbits per second. + - `penalty` (torch.Tensor): Commitment loss. + - `metrics` (dict): RVQ metrics, in particular rate of dead code replacement, and entropy. + """ + semantic_result = self.rvq_first(x, frame_rate) + if self.n_q == self.n_q_semantic: + return semantic_result + acoustic_result = self.rvq_rest(x, frame_rate) + full_quantized_emb = semantic_result.x + acoustic_result.x + full_quantized_codes = torch.cat( + [semantic_result.codes, acoustic_result.codes], dim=1 + ) + # This is the actual number of quantizers used, e.g. taking into account quantizer dropout. + n_q_semantic = semantic_result.codes.shape[1] + n_q_acoustic = acoustic_result.codes.shape[1] + full_quantized_bandwidth = semantic_result.bandwidth + acoustic_result.bandwidth + full_quantized_penalty = self._renorm_and_add( + semantic_result.penalty, acoustic_result.penalty, n_q_semantic, n_q_acoustic + ) + full_quantized_metrics = semantic_result.metrics + for key, value in acoustic_result.metrics.items(): + if key in full_quantized_metrics: + full_quantized_metrics[key] = self._renorm_and_add( + full_quantized_metrics[key], value, n_q_semantic, n_q_acoustic + ) + else: + full_quantized_metrics[key] = value + return QuantizedResult( + full_quantized_emb, + full_quantized_codes, + full_quantized_bandwidth, + penalty=full_quantized_penalty, + metrics=full_quantized_metrics, + ) + + def encode(self, x: torch.Tensor) -> torch.Tensor: + """Encode a given input tensor with the specified frame rate at the given bandwidth. + The RVQ encode method sets the appropriate number of quantizer to use + and returns indices for each quantizer. + """ + codes = self.rvq_first.encode(x) + if self.n_q > self.n_q_semantic: + acoustic_codes = self.rvq_rest.encode(x) + codes = torch.cat([codes, acoustic_codes], dim=1) + # codes is [B, K, T], with T frames, K nb of codebooks. + return codes + + def decode(self, codes: torch.Tensor) -> torch.Tensor: + """Decode the given codes to the quantized representation.""" + # codes is [B, K, T], with T frames, K nb of codebooks. + quantized = self.rvq_first.decode(codes[:, : self.n_q_semantic]) + if codes.shape[1] > self.n_q_semantic: + quantized += self.rvq_rest.decode(codes[:, self.n_q_semantic :]) + return quantized + + @property + def total_codebooks(self): + return self.rvq_first.max_n_q + self.rvq_rest.max_n_q + + @property + def num_codebooks(self): + return self.rvq_first.num_codebooks + self.rvq_rest.num_codebooks + + @property + def n_q(self): + return self.rvq_first.n_q + self.rvq_rest.n_q + + @property + def dimension(self): + return self.rvq_first.dimension + + @property + def semantic_quantizer(self) -> ResidualVectorQuantizer: + """This returns the quantizer that models the first level of the hierarchy (typically semantic).""" + return self.rvq_first + + @property + def acoustic_quantizer(self) -> ResidualVectorQuantizer: + """This returns the quantizer that models the higher levels of the hierarchy (typically acoustic).""" + return self.rvq_rest + + def set_num_codebooks(self, n: int): + assert n >= self.n_q_semantic and n <= self.total_codebooks + self.rvq_rest.set_num_codebooks(n - self.n_q_semantic) + + @property + def cardinality(self) -> int: + assert self.rvq_rest.cardinality == self.rvq_first.cardinality + return self.rvq_first.cardinality diff --git a/moshi_src/moshi/utils/__init__.py b/moshi_src/moshi/utils/__init__.py new file mode 100644 index 0000000..74bbb52 --- /dev/null +++ b/moshi_src/moshi/utils/__init__.py @@ -0,0 +1,10 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. +"""Utilities.""" diff --git a/moshi_src/moshi/utils/autocast.py b/moshi_src/moshi/utils/autocast.py new file mode 100644 index 0000000..e90efe0 --- /dev/null +++ b/moshi_src/moshi/utils/autocast.py @@ -0,0 +1,45 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +import torch + + +class TorchAutocast: + """TorchAutocast utility class. + Allows you to enable and disable autocast. This is specially useful + when dealing with different architectures and clusters with different + levels of support. + + Args: + enabled (bool): Whether to enable torch.autocast or not. + args: Additional args for torch.autocast. + kwargs: Additional kwargs for torch.autocast + """ + + def __init__(self, enabled: bool, *args, **kwargs): + self.autocast = torch.autocast(*args, **kwargs) if enabled else None + + def __enter__(self): + if self.autocast is None: + return + try: + self.autocast.__enter__() + except RuntimeError: + device = self.autocast.device + dtype = self.autocast.fast_dtype + raise RuntimeError( + f"There was an error autocasting with dtype={dtype} device={device}\n" + "If you are on the FAIR Cluster, you might need to use autocast_dtype=float16" + ) + + def __exit__(self, *args, **kwargs): + if self.autocast is None: + return + self.autocast.__exit__(*args, **kwargs) diff --git a/moshi_src/moshi/utils/compile.py b/moshi_src/moshi/utils/compile.py new file mode 100644 index 0000000..fe9c385 --- /dev/null +++ b/moshi_src/moshi/utils/compile.py @@ -0,0 +1,287 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +""" +Provides some extra utilities around torch compile, in particular with a way +to fully deactivate it easily with a context manager. +Provides a simple activation checkpointing that is compatible with FSDP and torch compile. +Finally, provides some utilities for CUDA graphing functions. +""" +from contextlib import contextmanager +from functools import wraps +import inspect +import os +import typing as tp + +import torch +from torch import cuda + + +_compile_disabled: bool = False + + +@contextmanager +def no_compile(): + """Disable torch.compile locally. Now Pytorch 2.4 provides a function to do that.""" + global _compile_disabled + + prev_disabled = _compile_disabled + _compile_disabled = True + try: + yield + finally: + _compile_disabled = prev_disabled + + +def torch_compile_lazy(fun): + """torch.compile creates a huge pool of processes, even when not using the function at all, + e.g. with Dora. This can polute stderr when doing CTRL+C. So we do it in a lazy way. + """ + if os.environ.get("NO_TORCH_COMPILE"): + return fun + fun_compiled = None + + @wraps(fun) + def _wrapped(*args, **kwargs): + nonlocal fun_compiled + if _compile_disabled: + return fun(*args, **kwargs) + if fun_compiled is None: + fun_compiled = torch.compile(fun) + return fun_compiled(*args, **kwargs) + + return _wrapped + + +class Checkpoint(torch.autograd.Function): + @staticmethod + def forward(ctx, function, *args) -> tp.Any: + to_save = [] + ctx.others = [] + ctx.function = function + # Sources will indicate whether the arg in position N is + # a tensor stored in ctx.save_for_backward, or inside ctx.others. + ctx.sources = [] + new_args = [] + for arg in args: + if isinstance(arg, torch.Tensor): + to_save.append(arg) + ctx.sources.append("tensor") + new_args.append(arg.detach()) + else: + ctx.sources.append("other") + ctx.others.append(arg) + new_args.append(arg) + ctx.save_for_backward(*to_save) + # During the forward, we just make a pass with no gradient computed. + with torch.no_grad(): + res = function(*new_args) + return res + + @staticmethod + def backward(ctx, *grads) -> tp.Tuple[tp.Optional[torch.Tensor], ...]: + pseudo_tensors = [] + with torch.set_grad_enabled(True): + # We create leaf tensors to collect the output gradients. + # We call them pseudo_tensors because they are pretending to be the input + # to `function` but are not directly + for tensor in ctx.saved_tensors: + pseudo_tensor = tensor.detach() + pseudo_tensor.requires_grad_(True) + pseudo_tensors.append(pseudo_tensor) + pseudo_tensors_copy = list(pseudo_tensors) + args = [] + for source in ctx.sources: + if source == "other": + args.append(ctx.others.pop(0)) + else: + assert source == "tensor" + args.append(pseudo_tensors_copy.pop(0)) + res = ctx.function(*args) + # The second forward with grad computation allows us to connect the input leaf tensors + # inside pseudo_tensors, to the outputs of the function called. + if not isinstance(res, tuple): + res = (res,) + # Now we just ask Torch to compute the derivative of `res` given the gradient coming from above + # `grads`. The computed gradient will end up into the `pseudo_tensors` grad attributes. + torch.autograd.backward(res, grads) + out: tp.List[tp.Optional[torch.Tensor]] = [None] + for source in ctx.sources: + # We still need to output `None` values for non tensor parameters. + if source == "other": + out.append(None) + else: + assert source == "tensor" + out.append(pseudo_tensors.pop(0).grad) + return tuple(out) + + +def simple_checkpoint(module: torch.nn.Module, *args, **kwargs): + """Custom implementation of checkpointing in PyTorch as the builtin implementation is broken + when using torch compile. Only supports wrapping a `nn.Module` with a forward with no `*args` or `**kwargs`. + + https://github.com/pytorch/pytorch/issues/97436. + Should be resolved in nightlies, but it is quite fun and simple to code it ourselves. + """ + if hasattr(module, "_fsdp_wrapped_module"): + module_for_sig = module._fsdp_wrapped_module + else: + module_for_sig = module + assert isinstance(module_for_sig, torch.nn.Module) + sig = inspect.signature(module_for_sig.forward) + # We first flatten all arguments to use only *args, to make things easier and because + # torch.autograd.Function has weird support for kwargs. + bounded = sig.bind(*args, **kwargs) + new_args = [] + for name, param in sig.parameters.items(): + if param.kind in { + inspect.Parameter.VAR_POSITIONAL, + inspect.Parameter.VAR_KEYWORD, + }: + raise RuntimeError("simple_checkpoint doesn't support var args.") + if name not in bounded.arguments: + break + new_args.append(bounded.arguments[name]) + return Checkpoint.apply(module, *new_args) + + +_in_cuda_graph = False +_disable_cuda_graph = False + + +def in_cuda_graph() -> bool: + """Indicate whether we are in a function that is CUDA Graphed (or will be soon).""" + return _in_cuda_graph + + +@contextmanager +def _set_in_cuda_graph(): + global _in_cuda_graph + assert not _in_cuda_graph + _in_cuda_graph = True + try: + yield + finally: + _in_cuda_graph = False + + +def _is_cuda_graph_enabled() -> bool: + if _disable_cuda_graph: + return False + no_cuda_graph = os.environ.get("NO_CUDA_GRAPH", "") + if no_cuda_graph.lower() not in {"0", "no", "n", ""}: + return False + return True + + +@contextmanager +def no_cuda_graph(): + """Deactivate CUDA Graphing for all the calls in this context manager.""" + global _disable_cuda_graph + old_value = _disable_cuda_graph + _disable_cuda_graph = True + try: + yield + finally: + _disable_cuda_graph = old_value + + +class CUDAGraphed: + """Allow simple CUDA Graphing of a function. + + Args: + func: callable, taking any number of arguments. Its tensors arguments should + be top level args, not nested in structures (tuples, dicts, etc). Keyword + arguments are NOT supported for simplicity. + warmup_steps: how many call to make normally before CUDA Graphing. In particular, this + allows torch.compiled functions to get properly compiled. + disabled: if True, just call the func directly, useful to quickly deactivate on CPU. + """ + + def __init__(self, func: tp.Callable, warmup_steps: int = 1, disable: bool = False): + self.func = func + self.warmup_steps = warmup_steps + self.disable = disable + self._graph: cuda.CUDAGraph | None = None + self._output: tuple | None = None + self._args: tuple | None = None + + def reset(self, warmup_steps: int = 0) -> None: + """Reset the state, meaning the next call we get CUDA Graphed again. Useful if some + shapes have changed, or external state (e.g. KVCache) has changed.""" + self.warmup_steps = warmup_steps + self._graph = None + self._output = None + self._args = None + + def __call__(self, *args, **kwargs) -> tp.Any: + if kwargs: + raise RuntimeError("Named arguments not supported for now.") + if self.disable or not _is_cuda_graph_enabled() or in_cuda_graph(): + return self.func(*args, **kwargs) + + def _clone_tensors(args: tuple) -> tuple: + out: list = [] + for arg in args: + if isinstance(arg, torch.Tensor): + arg = arg.clone() + out.append(arg) + return tuple(out) + + def _match_values_copy_tensors(args: tuple, target_args: tuple) -> None: + if len(args) != len(target_args): + raise ValueError( + f"Expected {len(target_args)}, but got {args} for CUDA Graphed function." + ) + for idx, (source, target) in enumerate(zip(args, target_args)): + if isinstance(target, torch.Tensor): + if not isinstance(source, torch.Tensor): + raise ValueError( + f"Argument #{idx} was a tensor, and is no longer (now {source})." + ) + if source.shape != target.shape: + raise ValueError( + f"Argument #{idx} had shape {target.shape}, but got shape {source.shape}" + "When using CUDAGraph, every call must be done with exactly the same shapes. " + "Feel free to deactivate with the env variable NO_CUDA_GRAPH=1, or the decorator " + "`with no_cuda_graph():`" + ) + target.copy_(source) + else: + if isinstance(source, torch.Tensor): + raise ValueError( + f"Argument #{idx} was not a tensor {target}, but is now one." + ) + if source is not target and source != target: + raise ValueError( + f"Argument #{idx} changed value from {target} to {source}." + ) + + with _set_in_cuda_graph(): + # Prevent any one under us to try and CUDA Graph things. + if self._graph is None: + if self.warmup_steps <= 0: + self._graph = cuda.CUDAGraph() + # Making a copy just to ensure those are not used else where. + self._args = _clone_tensors(args) + with cuda.graph(self._graph): + self._output = self.func(*self._args) + # At this point nothing really happened, so we have to make it run for real. + self._graph.replay() + return self._output + else: + self.warmup_steps -= 1 + return self.func(*args) + else: + assert self._args is not None + _match_values_copy_tensors(args, self._args) + self._graph.replay() + return self._output + + +def cuda_graph(func: tp.Callable, warmup_steps: int = 1): + """Just calls `CUDAGraphed` on the given function.""" + if not _is_cuda_graph_enabled(): + return func + return CUDAGraphed(func, warmup_steps) diff --git a/moshi_src/moshi/utils/quantize.py b/moshi_src/moshi/utils/quantize.py new file mode 100644 index 0000000..70bb3b6 --- /dev/null +++ b/moshi_src/moshi/utils/quantize.py @@ -0,0 +1,57 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +"""Quantization based on bitsandbytes, supporting only 8 bits for now. +We are taking from freedom from the intended use of bnb: +""" + +import torch +from torch import nn + + +class QLinear(nn.Module): + def __init__(self, linear: nn.Linear): + super().__init__() + from bitsandbytes import functional as bnbF # type: ignore + weight = linear.weight + assert weight.data.dtype.is_floating_point + assert linear.bias is None + CB, SCB, _ = bnbF.int8_vectorwise_quant(weight.data.to(torch.float16)) # type: ignore + self.weight = nn.Parameter(CB, requires_grad=False) + self.weight_scb = nn.Parameter(SCB, requires_grad=False) + + def forward(self, x): + import bitsandbytes as bnb # type: ignore + state = bnb.MatmulLtState() + state.CB = self.weight # type: ignore + assert isinstance(state.CB, torch.Tensor) + state.SCB = self.weight_scb # type: ignore + assert isinstance(state.SCB, torch.Tensor) + if state.SCB.dtype != torch.float: + raise RuntimeError( + "Expected `weight_scb` to have type float, but got bfloat16. " + "When using quantized models, care should be taken not to change the dtype of " + "the model once initialized.") + assert state.SCB.dtype == torch.float, state.SCB.dtype + state.has_fp16_weights = False + y = bnb.matmul(x.half(), state.CB, state=state) + assert isinstance(y, torch.Tensor) + return y + + +def replace_linear_with_qlinear(module): + """Recursively replace all Linear layers with QLinear layers.""" + for name, child in module.named_children(): + if isinstance(child, nn.Linear): + setattr(module, name, QLinear(child)) + elif isinstance(child, QLinear): + # Slight issue with the way we implement things: the scale param + # might get casted with the rest of the model to bfloat16, altough + # we most likely want to keep it as float. For the LM model we might call this function twice, + # first layer by layer to avoid to big of a memory usage, and second, at the end + # of the LM init, after all other modules are initialized and properly dtyped. + # In any case that should happen before loading the state dict to avoid a loss of precision. + child.float() + else: + replace_linear_with_qlinear(child) diff --git a/moshi_src/moshi/utils/sampling.py b/moshi_src/moshi/utils/sampling.py new file mode 100644 index 0000000..2a21d0a --- /dev/null +++ b/moshi_src/moshi/utils/sampling.py @@ -0,0 +1,127 @@ +# Copyright (c) Kyutai, all rights reserved. +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the license found in the +# LICENSE file in the root directory of this source tree. + + +import torch + + +def multinomial( + input: torch.Tensor, num_samples: int, replacement=False, *, generator=None +): + """torch.multinomial with arbitrary number of dimensions, and number of candidates on the last dimension. + + Args: + input (torch.Tensor): The input tensor containing probabilities. + num_samples (int): Number of samples to draw. + replacement (bool): Whether to draw with replacement or not. + Keywords args: + generator (torch.Generator): A pseudorandom number generator for sampling. + Returns: + torch.Tensor: Last dimension contains num_samples indices + sampled from the multinomial probability distribution + located in the last dimension of tensor input. + """ + input_ = input.reshape(-1, input.shape[-1]) + # We should probably be able to remove this once the following PR has landed: + # https://github.com/pytorch/pytorch/pull/134818/files + # In the meantime, we specialize the case no-replacement, nsamples=1 so as not + # to have a synchronization point. + if replacement or num_samples != 1: + output_ = torch.multinomial( + input_, + num_samples=num_samples, + replacement=replacement, + generator=generator, + ) + else: + q = torch.empty_like(input_).exponential_(1, generator=generator) + q = input_ / q + output_ = q.argmax(dim=-1, keepdim=True) + output = output_.reshape(*list(input.shape[:-1]), -1) + return output + + +def sample_top_k(probs: torch.Tensor, k: int) -> torch.Tensor: + """Sample next token from top K values along the last dimension of the input probs tensor. + + Args: + probs (torch.Tensor): Input probabilities with token candidates on the last dimension. + k (int): The k in “top-k”. + Returns: + torch.Tensor: Sampled tokens. + """ + k = min(k, probs.shape[-1]) + probs, indices = torch.topk(probs, k, dim=-1) + next_token = multinomial(probs, num_samples=1) + next_token = indices.gather(-1, next_token) + return next_token + + +def sample_top_p(probs: torch.Tensor, p: float) -> torch.Tensor: + """Sample next token from top P probabilities along the last dimension of the input probs tensor. + + Args: + probs (torch.Tensor): Input probabilities with token candidates on the last dimension. + p (int): The p in “top-p”. + Returns: + torch.Tensor: Sampled tokens. + """ + probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True) + probs_sum = torch.cumsum(probs_sort, dim=-1) + mask = probs_sum - probs_sort > p + probs_sort *= (~mask).float() + probs_sort.div_(probs_sort.sum(dim=-1, keepdim=True)) + next_token = multinomial(probs_sort, num_samples=1) + next_token = torch.gather(probs_idx, -1, next_token) + return next_token + + +def sample_token( + logits: torch.Tensor, + use_sampling: bool = False, + temp: float = 1.0, + top_k: int = 0, + top_p: float = 0.0, +) -> torch.Tensor: + """Given logits of shape [*, Card], returns a LongTensor of shape [*].""" + # Apply softmax for sampling if temp > 0. Else, do greedy sampling to avoid zero division error. + if use_sampling and temp > 0.0: + probs = torch.softmax(logits / temp, dim=-1) + if top_p > 0.0: + next_token = sample_top_p(probs, p=top_p) + elif top_k > 0: + next_token = sample_top_k(probs, k=top_k) + else: + next_token = multinomial(probs, num_samples=1) + else: + next_token = torch.argmax(logits, dim=-1, keepdim=True) + assert next_token.shape[-1] == 1 + return next_token[..., 0] + + +if __name__ == "__main__": + torch.manual_seed(1234) + device = "cpu" + if torch.cuda.is_available(): + torch.backends.cuda.matmul.allow_tf32 = False + torch.backends.cudnn.allow_tf32 = False + device = "cuda:0" + + ps = torch.tensor([5.0, 2.0, 12.0, 6.0, 8.0, 1.0, 0.0, 4.0], device=device) + cnts = torch.zeros(ps.shape, dtype=torch.long, device=device) + total_samples = 1000 + for _ in range(total_samples): + vs = multinomial(ps, num_samples=1, replacement=False) + cnts[vs] += 1 + diff = cnts / cnts.sum() - ps / ps.sum() + max_diff = diff.abs().max().cpu().item() + print(ps / ps.sum()) + print(cnts / cnts.sum()) + assert max_diff < 1.5e-2 diff --git a/moshi_src/moshi/utils/utils.py b/moshi_src/moshi/utils/utils.py new file mode 100644 index 0000000..51eea20 --- /dev/null +++ b/moshi_src/moshi/utils/utils.py @@ -0,0 +1,52 @@ +import torch + +from .compile import torch_compile_lazy + + +@torch_compile_lazy +def cross_entropy( + logits: torch.Tensor, targets: torch.Tensor, mask: torch.Tensor, dtype=torch.float32, + logits_soft_clip: float | None = None) -> torch.Tensor: + """Compute cross entropy between multi-codebook targets and model's logits. + The cross entropy is computed per codebook to provide codebook-level cross entropy. + Valid timesteps for each of the codebook are pulled from the mask, where invalid + timesteps are set to 0. + + Args: + logits (torch.Tensor): Model's logits of shape [B, K, T, card]. + targets (torch.Tensor): Target codes, of shape [B, K, T]. + mask (torch.Tensor): Mask for valid target codes, of shape [B, K, T]. + dtype (type): Data type of the output cross entropy. + logits_soft_clip (float): Clipping value for the logits to avoid numerical instability. + Recommended value: 30.0. + Returns: + ce (torch.Tensor): Cross entropy [B, K, T] with type dtype. + """ + output_shape = targets.shape + assert logits.shape[:-1] == targets.shape + assert mask.shape == targets.shape + logits = logits.view(-1, logits.shape[-1]) + targets = targets.reshape(-1) + mask = mask.reshape(-1) + + safe_targets = torch.where( + mask, + targets, + torch.zeros(1, device=targets.device, dtype=targets.dtype), + ) + + # Chunking the conversion to float32 to avoid OOMs. + ce_chunks = [] + for logits_chunk, targets_chunk in zip(torch.chunk(logits, 4), torch.chunk(safe_targets, 4)): + logits_chunk = logits_chunk.to(dtype) + if logits_soft_clip is not None: + logits_chunk = logits_soft_clip * torch.tanh(logits_chunk / logits_soft_clip) + log_partition = torch.logsumexp(logits_chunk, dim=-1, keepdim=True) + + # For some reason, the PyTorch cross entropy is super slow with inputs with large cardinality (e.g. 32000) + # so we reimplement the cross entropy ourselves... + ce_chunks.append(log_partition - logits_chunk.gather(-1, targets_chunk[..., None])) + ce = torch.cat(ce_chunks, dim=0) + ce = ce[..., 0] + ce = torch.where(mask, ce, torch.zeros(1, device=ce.device, dtype=ce.dtype)) + return ce.view(output_shape) diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..b6a6f74 --- /dev/null +++ b/nodes.py @@ -0,0 +1,184 @@ +import torch +import torch._dynamo +import sys +import os + + +os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'garbage_collection_threshold:0.1' +import folder_paths +from pathlib import Path +import json +import random +import numpy as np +import comfy.utils +from tqdm import tqdm + +# Add the correct moshi source directory to the Python path +sys.path.insert(0, os.path.join(os.path.dirname(__file__), "moshi_src")) + +from moshi.models.loaders import CheckpointInfo +from moshi.models.tts import TTSModel + +# Monkey-patch the problematic function in the moshi library +# This prevents a PyTorch compilation error on Windows by disabling +# the JIT compiler for this specific function. +try: + import moshi.modules.rope + torch._dynamo.disable(moshi.modules.rope.apply_rope) +except (ImportError, AttributeError) as e: + print(f"KyutaiTTS Node: Could not patch moshi.modules.rope.apply_rope. If you encounter an OverflowError, this may be the cause. Error: {e}") + +class KyutaiTTS: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "text": ("STRING", {"multiline": True, "default": "Hey there! How are you?"}), + "model_path": ("STRING", {"default": "", "multiline": False, "folder_input": True}), + "voice_model": (folder_paths.get_filename_list("loras"), ), + "device": (["cuda", "cpu"],), + "n_q": ("INT", {"default": 32}), + "temp": ("FLOAT", {"default": 0.6, "min": 0.0, "max": 1.0, "step": 0.1}), + "cfg_coef": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 10.0, "step": 0.1}), + "padding_between": ("INT", {"default": 1}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFF}), + } + } + + RETURN_TYPES = ("AUDIO", ) + FUNCTION = "generate" + CATEGORY = "Kyutai" + + def generate(self, text, model_path, voice_model, device, n_q, temp, cfg_coef, padding_between, seed): + + + def seed_all(seed): + torch.manual_seed(seed) + random.seed(seed) + np.random.seed(seed) + + seed_all(seed) + + device = torch.device(device) + + # Create CheckpointInfo from the local model path + full_model_path = model_path + if not full_model_path or not os.path.isdir(full_model_path): + raise FileNotFoundError(f"Model directory not found: {full_model_path}") + + # Define expected file names within the model directory + moshi_weights_path = os.path.join(full_model_path, "dsm_tts_1e68beda@240.safetensors") + if not os.path.exists(moshi_weights_path): + raise FileNotFoundError(f"Moshi weights (dsm_tts_1e68beda@240.safetensors) not found in {full_model_path}") + + mimi_weights_path = os.path.join(full_model_path, "tokenizer-e351c8d8-checkpoint125.safetensors") + if not os.path.exists(mimi_weights_path): + raise FileNotFoundError(f"Mimi weights (tokenizer-e351c8d8-checkpoint125.safetensors) not found in {full_model_path}") + + tokenizer_path = os.path.join(full_model_path, "tokenizer_spm_8k_en_fr_audio.model") + if not os.path.exists(tokenizer_path): + raise FileNotFoundError(f"Tokenizer (tokenizer_spm_8k_en_fr_audio.model) not found in {full_model_path}") + + config_path = os.path.join(full_model_path, "config.json") + if not os.path.exists(config_path): + raise FileNotFoundError(f"config.json not found in {full_model_path}") + + with open(config_path, 'r') as f: + raw_config = json.load(f) + + # Extract specific configs for CheckpointInfo and remove them from lm_config + tts_config = raw_config.get("tts_config", {}) + stt_config = raw_config.get("stt_config", {}) + lm_gen_config = raw_config.get("lm_gen_config", {}) + model_id = raw_config.get("model_id", {}) + model_type = raw_config.get("model_type", "moshi") # Extract model_type + + lm_config = dict(raw_config) # Create a copy for lm_config + + # Remove keys not meant for LMModel from lm_config + lm_config.pop("tts_config", None) + lm_config.pop("stt_config", None) + lm_config.pop("lm_gen_config", None) + lm_config.pop("model_id", None) + lm_config.pop("moshi_name", None) + lm_config.pop("mimi_name", None) + lm_config.pop("tokenizer_name", None) + lm_config.pop("lora_name", None) + lm_config.pop("model_type", None) # Remove model_type from lm_config + + checkpoint_info = CheckpointInfo( + moshi_weights=Path(moshi_weights_path), + mimi_weights=Path(mimi_weights_path), + tokenizer=Path(tokenizer_path), + lm_config=lm_config, + raw_config=raw_config, + tts_config=tts_config, + stt_config=stt_config, + lm_gen_config=lm_gen_config, + model_id=model_id, + model_type=model_type # Pass model_type to CheckpointInfo + ) + + tts_model = TTSModel.from_checkpoint_info( + checkpoint_info, n_q=n_q, temp=temp, device=device + ) + + entries = tts_model.prepare_script([text], padding_between=padding_between) + + voice_path = folder_paths.get_full_path("loras", voice_model) + if not voice_path or not os.path.exists(voice_path): + raise FileNotFoundError(f"Voice model not found: {voice_model}") + + condition_attributes = tts_model.make_condition_attributes( + [voice_path], cfg_coef=cfg_coef + ) + + # --- Step 1: Generate audio token frames --- + frames_list = [] + # A more accurate estimation including initial/final padding, model delays, and an empirical fudge factor. + initial_padding = tts_model.machine.initial_padding + final_padding = tts_model.final_padding + delay_steps = tts_model.delay_steps + # Add a small fudge factor for each word to account for un-predictable discretionary padding. + FUDGE_FACTOR_PER_WORD = 1 + word_steps = sum(len(entry.tokens) + entry.padding + FUDGE_FACTOR_PER_WORD for entry in entries) + total_steps = initial_padding + word_steps + delay_steps + final_padding + gen_pbar = comfy.utils.ProgressBar(total_steps) + with tqdm(total=total_steps, desc="Generating Tokens") as pbar_cmd_gen: + def _on_frame_collect(frame): + if (frame != -1).all(): + frames_list.append(frame.clone()) + # Update by 1 for each frame generated. + gen_pbar.update(1) + pbar_cmd_gen.update(1) + + all_entries = [entries] + all_condition_attributes = [condition_attributes] + with tts_model.mimi.streaming(len(all_entries)): + tts_model.generate(all_entries, all_condition_attributes, on_frame=_on_frame_collect) + + # --- Step 2: Decode frames to PCM audio --- + pcms = [] + if frames_list: + decode_pbar = comfy.utils.ProgressBar(len(frames_list)) + with tqdm(total=len(frames_list), desc="Decoding Audio") as pbar_cmd_decode: + for frame in frames_list: + pcm = tts_model.mimi.decode(frame[:, 1:, :]).cpu().numpy() + pcms.append(np.clip(pcm[0], -1, 1)[np.newaxis, :]) + decode_pbar.update(1) + pbar_cmd_decode.update(1) + + # --- Step 3: Concatenate audio chunks --- + audio = np.concatenate(pcms, axis=-1) if pcms else np.array([]) + + print(f"KyutaiTTS Node: Outputting audio with sample rate: {tts_model.mimi.sample_rate}") + # Return audio in the format expected by ComfyUI's AUDIO type + return ({"waveform": torch.from_numpy(audio), "sample_rate": tts_model.mimi.sample_rate},) + +NODE_CLASS_MAPPINGS = { + "KyutaiTTS": KyutaiTTS, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "KyutaiTTS": "KyutaiTTS", +} \ No newline at end of file diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..7ce057e --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +sphn<0.2 \ No newline at end of file