diff --git a/.gitattributes b/.gitattributes index 352207e9d567da48dbf850b4581f78a88e8f72b4..0d6fd30a5b4671382ed4280e16d66e0428156cee 100644 --- a/.gitattributes +++ b/.gitattributes @@ -45,4 +45,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text *.wasm filter=lfs diff=lfs merge=lfs -text *.zst filter=lfs diff=lfs merge=lfs -text *tfevents* filter=lfs diff=lfs merge=lfs -text -"/model.safetensors.index.json" filter=lfs diff=lfs merge=lfs -text \ No newline at end of file +"/model.safetensors.index.json" filter=lfs diff=lfs merge=lfs -textmodel.safetensors.index.json filter=lfs diff=lfs merge=lfs -text diff --git a/model-00233-of-00341.safetensors b/model-00233-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..3efa87ca88f45e74b9e1862954fb6e8cf7be9cda --- /dev/null +++ b/model-00233-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d0d06bb4db7f611a4d93397374b3c1e65b2fe4c9c31eb9254e1886a1257beaa8 +size 3152265592 diff --git a/model-00234-of-00341.safetensors b/model-00234-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..5636df01e6be0b0a363392d6576c849b659e28d1 --- /dev/null +++ b/model-00234-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:820d949eb646d381b488f66e1faa238d9cfcf0729048ac57b95c4cb48cd7c5cb +size 3153288400 diff --git a/model-00235-of-00341.safetensors b/model-00235-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..76db4952c0cfbe861359a857068c3704e669d88a --- /dev/null +++ b/model-00235-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e4278badf10b3faee48796ea68964828fd2097da38250168a758f36a570cdc7d +size 2401529096 diff --git a/model-00236-of-00341.safetensors b/model-00236-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..9e5e8076b2fdef7d3e45813998b92416018d2374 --- /dev/null +++ b/model-00236-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5be50ff6290c2988985d4219dc25398a7e1abe7be487a3b9461139836edfc080 +size 3154458912 diff --git a/model-00237-of-00341.safetensors b/model-00237-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..e5d30c12cbeb257d58e6a7c787438b63d2df78fc --- /dev/null +++ b/model-00237-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4d1fe5ac4552381e444d5427075ba3430b7d8a1acad416288fcb93c11e153698 +size 3153288488 diff --git a/model-00238-of-00341.safetensors b/model-00238-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..a440d8bae81697377c25480ea35e6eb38c2f3be1 --- /dev/null +++ b/model-00238-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a3fda316f38b9cc4c4ef5b9939bc8ecd7795f5c08b7e314ca755ee54f673bea0 +size 2822748248 diff --git a/model-00239-of-00341.safetensors b/model-00239-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..603848a0e11d3074bedc563cce0118f3df823039 --- /dev/null +++ b/model-00239-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7c41579aa8ae4b089a6221bd8952bd33f2193754f9abd1f818cda1721f2eb930 +size 3154458912 diff --git a/model-00240-of-00341.safetensors b/model-00240-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..3c6abddc95f27dd3c93b8d21eb9393058160c576 --- /dev/null +++ b/model-00240-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c3f434ac1460219e93abaecbe3b179521481c614d6a37a8160df8e621907b7ae +size 3153288488 diff --git a/model-00241-of-00341.safetensors b/model-00241-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..919b9dae9cdf593ce97f1dbfa5da44b5f5aa7a90 --- /dev/null +++ b/model-00241-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:50486fafb712caff18616daf0c76adf609dbac00d9a506bf0529013b148e6925 +size 2822748248 diff --git a/model-00242-of-00341.safetensors b/model-00242-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..2c7509806e11d83541ae07266d1f5a3e0c368c4a --- /dev/null +++ b/model-00242-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c110ca86ca516846b672ca6108adfafa1c8d8e6cbd090df2909cd37a3bb312ae +size 3154458912 diff --git a/model-00243-of-00341.safetensors b/model-00243-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..1565f6aba5a1971223845af8e46d89b137697012 --- /dev/null +++ b/model-00243-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b97ce558ff7caed82d2af08acb154aff1b597ccf2c99be8ea5454972833e54e4 +size 3153288488 diff --git a/model-00244-of-00341.safetensors b/model-00244-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..3990d77683d83fe872d158a40181ef85d8896469 --- /dev/null +++ b/model-00244-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:370990aa6144562709ecc7742fa493342e1051427453ba6d7cce0087a8a20c41 +size 2822748248 diff --git a/model-00245-of-00341.safetensors b/model-00245-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..6d21935f49f98d16e070aa1d43b105f373980801 --- /dev/null +++ b/model-00245-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0a5bdf2b1938c8070e46b4db6e7740a1d6d17ddb313c6e42367134399c6d3697 +size 3152265592 diff --git a/model-00246-of-00341.safetensors b/model-00246-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..272f23010722421484653a4feea6758445a39d42 --- /dev/null +++ b/model-00246-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:535a6cc52b030ba2038421538e071b3f38da1c25eebf9a3534aff3fda2852783 +size 3153288400 diff --git a/model-00247-of-00341.safetensors b/model-00247-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..d552389da957ba23abc18f97df7a1b7ea1ea13b0 --- /dev/null +++ b/model-00247-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4ec2014601731047e705080f513d7ef7f0868514f58af3460ae3035c1f3eb01f +size 2401529096 diff --git a/model-00248-of-00341.safetensors b/model-00248-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..4ad7e71b6545bd318ccc92482adf15c65332cf97 --- /dev/null +++ b/model-00248-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:356dc70925fc018d7fe29590da9d037a951abdeae4afc1017f226b3346b3500d +size 3154458912 diff --git a/model-00249-of-00341.safetensors b/model-00249-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..e8f33430f04527b154c0af76e39ecfdcdada3916 --- /dev/null +++ b/model-00249-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3c5e73119012d3fc0b1545618b3da4c5c149673a394b2cfc85335b9e4ce25dd7 +size 3153288488 diff --git a/model-00250-of-00341.safetensors b/model-00250-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..a7bf87dd6756a24989a0fd127961d818d1106a7a --- /dev/null +++ b/model-00250-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:665bc4b798c34661af79f4019c98742df11d37861a0d9deec34b13ab7cba2f21 +size 2822748248 diff --git a/model-00251-of-00341.safetensors b/model-00251-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..0c0874ebe35b5107cac133496d5bd82477124354 --- /dev/null +++ b/model-00251-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bf4a68175d7865be3d4d79ac801f3e9224f2f7149c883b68819faf88db4a43e0 +size 3154458912 diff --git a/model-00252-of-00341.safetensors b/model-00252-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..3f0adee1b9fcbd51506ce742957083ba7cfa6562 --- /dev/null +++ b/model-00252-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:de3a9108a3bd3489b6912e7cfb75cc51d4188951511eac40e436c112fe6048f0 +size 3153288488 diff --git a/model-00253-of-00341.safetensors b/model-00253-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..5ac4e2fa0b9ece913db6d0e3958de25c437239ef --- /dev/null +++ b/model-00253-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b20ee9ad5fdf76aae64b0b2d62bdef1ecf4db251677aba47528abb4d8a9cf3e9 +size 2822748248 diff --git a/model-00254-of-00341.safetensors b/model-00254-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c08d81fded51f93a15a52e2df55aad9409e931b2 --- /dev/null +++ b/model-00254-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:397f707ce9402e901a63ed2cfb575cf4a3b45d76825191a03bf052e6a85c14cc +size 3154458912 diff --git a/model-00255-of-00341.safetensors b/model-00255-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..156182477d872f065d9f68b8f808010409487e9a --- /dev/null +++ b/model-00255-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a9a647f832aeb4c3ffd6a56e5790c84e4fd8173ae5978b5c9a818af7a8adbd3e +size 3153288488 diff --git a/model-00256-of-00341.safetensors b/model-00256-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..ed17ed61bf0142f93a21514a925e54a6c0ee8766 --- /dev/null +++ b/model-00256-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:46685e741a3d93d8fe0300b2aae7f2cfa21aad63efef75bc073c6875128d752d +size 2822748248 diff --git a/model-00257-of-00341.safetensors b/model-00257-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..cd395cbe1bfe0c6adad3e0c3ce05bc80018ca82c --- /dev/null +++ b/model-00257-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:840de9b3675a1809a632255ee8147148f5191b9c540ceb9b3602e2443f3ec509 +size 3152265592 diff --git a/model-00258-of-00341.safetensors b/model-00258-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..a7b1fd5c6f0382eeed3c82492498892cb561582e --- /dev/null +++ b/model-00258-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e5a2d98704a3b35b2fcd7ff3e6edd13f6581440da616b69d3505fc3049e27fbc +size 3153288400 diff --git a/model-00259-of-00341.safetensors b/model-00259-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..47a5245b4efda8a12b4b42118c7ad24914b13fce --- /dev/null +++ b/model-00259-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0200e5b3a3c3b0ab80fc3e29f8618c9d1d7fd7629b2afadcae1ed5f7b6b51bd9 +size 2401529096 diff --git a/model-00260-of-00341.safetensors b/model-00260-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..7096fa97635e6d964ca073d080412ff182bbbaad --- /dev/null +++ b/model-00260-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d353a7f5cd3c1f5795e6a71d9a21b4026661a3d014f8dc6d1dc13ff349c5283e +size 3154458912 diff --git a/model-00261-of-00341.safetensors b/model-00261-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..9b797dfebc0eadab16091284cd0c0f5e0b919812 --- /dev/null +++ b/model-00261-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3a4ef454196f63efd86b88042ddc44a785d353fc71be43e0b03b5cadb872f338 +size 3153288488 diff --git a/model-00262-of-00341.safetensors b/model-00262-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..e30e224eb9ddb7cced06ce01ea8feade69f2253c --- /dev/null +++ b/model-00262-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09b8fabb607ac21d289583ee6f6ed4970528153eebb5de37f37896bf77ac8b47 +size 2822748248 diff --git a/model-00263-of-00341.safetensors b/model-00263-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..45765e19a3c67a3e7abbe92187fe000d3b7420ec --- /dev/null +++ b/model-00263-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:921688f24e3d67c704ae93b224da32aa3daac7a3fb5d29203e14895cb5cf1e3a +size 3154458912 diff --git a/model-00264-of-00341.safetensors b/model-00264-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..83a32ff70c7b3f29ab4831bdeb0a79ce1b2abbd7 --- /dev/null +++ b/model-00264-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97d556848fdd8c93c9afec6b07c057e6813d8f5bad2b007ebcabfd0b6c5c16d4 +size 3153288488 diff --git a/model-00265-of-00341.safetensors b/model-00265-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..2fc5de28ae26d5d7da83d59706f18077af3ddd45 --- /dev/null +++ b/model-00265-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7eea1c6efe7647de9d1054047ea14e94a624fc4d4f61008e010e6d90524f8ae2 +size 2822748248 diff --git a/model-00266-of-00341.safetensors b/model-00266-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..46b80b2fc635b7d6f66673f15d965639549b6652 --- /dev/null +++ b/model-00266-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:948993a654e3f48aaf520e8b23289d13a91c323801d8b36027545b4eeefbd4a3 +size 3154458912 diff --git a/model-00267-of-00341.safetensors b/model-00267-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..8d9a26334314a9406503ea397988d00da61b3843 --- /dev/null +++ b/model-00267-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d27f4b0207ee84ad9ad7072df82b439b4e9d97bdb78ce60e9290f76921f44858 +size 3153288488 diff --git a/model-00268-of-00341.safetensors b/model-00268-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..d4c00dc9dd14208ea2cc3dd4a435e3fc1bffd04c --- /dev/null +++ b/model-00268-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5bf1e600edfe05e60f977f3bc6d1e0ee144988b9dbbc93bdfa027f4c6c6a6f8b +size 2822748248 diff --git a/model-00269-of-00341.safetensors b/model-00269-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..6a65a68033134dda913f2bffc7783fc73f23de67 --- /dev/null +++ b/model-00269-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:633eadc62b39434780f579acc4e9119cad8a495c515cf413dff00f45eb90e47b +size 3152265592 diff --git a/model-00270-of-00341.safetensors b/model-00270-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..115c657b01ffea12ae77ece7e175d46716832986 --- /dev/null +++ b/model-00270-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c17bc7afba6a80eecd90082a9b8c17cf85b1b05b0fa759ab14edd59ff8802bf4 +size 3153288400 diff --git a/model-00271-of-00341.safetensors b/model-00271-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..48e3aa4044dbb4ee2900f00eeebf6a36d83bc2ed --- /dev/null +++ b/model-00271-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:37866bd9bf8ee01c2135f8b04ec841b6891dafd96e4f4537b5c40724993cb5ee +size 2401529096 diff --git a/model-00272-of-00341.safetensors b/model-00272-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..062cb5c496e05bb43efe4ad0a38de0ddc6138ef4 --- /dev/null +++ b/model-00272-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:149f1e12e5db1ceb30d633a1241b374196b2f970ee80c561d79ab5e092eb4ba5 +size 3154458912 diff --git a/model-00273-of-00341.safetensors b/model-00273-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..a74998c9c1476e65c27d20ad2eea38276eeccf1a --- /dev/null +++ b/model-00273-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bc78bbf5ff572ce288a961014eb2b5c5280edc899cd8b73be9c073cf6892016c +size 3153288488 diff --git a/model-00274-of-00341.safetensors b/model-00274-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..f19918953c62bfceb65375b578e75797386e0b83 --- /dev/null +++ b/model-00274-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e93de251c5561424cef86832cbce575f63dd5d9b059db6ee783428301fc5db0a +size 2822748248 diff --git a/model-00275-of-00341.safetensors b/model-00275-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..6a5ea2a646b0e29062d39fd92b7343cf8bd7ee6c --- /dev/null +++ b/model-00275-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fe9d87636bafcc7780846966ed68cb447de1b5940ade2c250b8daec8d8d8cfc9 +size 3154458912 diff --git a/model-00276-of-00341.safetensors b/model-00276-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..b3ef3890185cbc5be243a51f55c273aae800bdbf --- /dev/null +++ b/model-00276-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f40e1044f2b89c67d3d971478874939f1c786ada49c9c832f71487f48bb95e23 +size 3153288488 diff --git a/model-00277-of-00341.safetensors b/model-00277-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..183690a611fac4cbf411a848c60d0c90ea13b817 --- /dev/null +++ b/model-00277-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3226afe7480c300962c373548db038ad655cee5fe547733662a21c187fc40a0d +size 2822748248 diff --git a/model-00278-of-00341.safetensors b/model-00278-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..07254042d6d3fcb8a7364a3b77888278ac0846b2 --- /dev/null +++ b/model-00278-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3636beec28b8e2338098fa45dfbd0593d20c5a7c8c754e9888fec05b22132e43 +size 3154458912 diff --git a/model-00279-of-00341.safetensors b/model-00279-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..aeb55ae1b8ec0b6d453fbc7ecd425de2dcadfa9a --- /dev/null +++ b/model-00279-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7f01ea2cd959c3471bda079c25f1012353987b848dbdee7ccb0b83d68e054c8a +size 3153288488 diff --git a/model-00280-of-00341.safetensors b/model-00280-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..f43ccba056ca264570ab7652d0e4749336a94f08 --- /dev/null +++ b/model-00280-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d36e6d6ab4633c5f6d62c998433c339a7149d4439daeef6540515517ed4d1ff3 +size 2822748248 diff --git a/model-00281-of-00341.safetensors b/model-00281-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..51d9c8fa2ea3ba1f2841747197aaaf83a20ba35d --- /dev/null +++ b/model-00281-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:97b6e17a50c7763cb6147fe6cfff57ecf163afe91a074bfad87749c10aef819c +size 3152265592 diff --git a/model-00282-of-00341.safetensors b/model-00282-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..11ffccab93e2aa174ddd64a2b8f6dcc52b4901d7 --- /dev/null +++ b/model-00282-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f314c4dbcd7cbe1b9a387c89d5a5c2ad51c5eb15f898026c88a82258bc98f9bf +size 3153288400 diff --git a/model-00283-of-00341.safetensors b/model-00283-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..bfc95f5a6dc55b9a9be9a966a2113755e8f4fcc7 --- /dev/null +++ b/model-00283-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:09c2619a0d92ad45768517bf187e3f3353de20b7477e2330eb1706f7d8a8c4be +size 2401529096 diff --git a/model-00284-of-00341.safetensors b/model-00284-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..1807cd588b241d3cb2457d437ff24c397372cd92 --- /dev/null +++ b/model-00284-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f86efd380325f2b515e427bf6b7777fe5dd9ed19296021ef8c0c9bb8d2518850 +size 3154458912 diff --git a/model-00285-of-00341.safetensors b/model-00285-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..deba98957752c2e4c251d66fb338210ea2ffdf57 --- /dev/null +++ b/model-00285-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9ff4bb9322593ca6624138414b6c023f18b11d26558e51a9e1c65db6de75c644 +size 3153288488 diff --git a/model-00286-of-00341.safetensors b/model-00286-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..6820954124ccd2850a538a8b5c219662cf511898 --- /dev/null +++ b/model-00286-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3049f2ab9a4f235f3c39c88b2a8a0bdbe044880b2a3f00acacd7a78176b56374 +size 2822748248 diff --git a/model-00287-of-00341.safetensors b/model-00287-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..8de28d9945c63ffd0aca3dccb4524b1b32e87c44 --- /dev/null +++ b/model-00287-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8e3eba5b6c73a680d96d3d9844962141a5ce08ee6c57c2a410da4205d4525093 +size 3154458912 diff --git a/model-00288-of-00341.safetensors b/model-00288-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..037db36cd315e4df00ae4e5ce38fe968af71660c --- /dev/null +++ b/model-00288-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c5295f7612fdaaccee96c354cd2508b28d74f317dd00716bf2471ec0988ae7f8 +size 3153288488 diff --git a/model-00289-of-00341.safetensors b/model-00289-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..eab5284a3e83baa71b428ada325e4fe188f54a81 --- /dev/null +++ b/model-00289-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e184cbcf5b7451554ce9e261fdb46ddb65472f401b1b07c2dbe6a8e98deb6f9f +size 2822748248 diff --git a/model-00290-of-00341.safetensors b/model-00290-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..8721106e71b7096c39d188bac8a4ea539733e18f --- /dev/null +++ b/model-00290-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8ee4f9b835299ea4c28e6cf35dd77be66e8325e1b872c47b90af0818f2c121f5 +size 3154458912 diff --git a/model-00291-of-00341.safetensors b/model-00291-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..b98561fcc71ff0e61a309b2dca2fae0ae60a5cb2 --- /dev/null +++ b/model-00291-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3128944c86c0ec9a7d604bdae948abd32e12f441d467b6da58925d7740b2a15 +size 3153288488 diff --git a/model-00292-of-00341.safetensors b/model-00292-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..3ee1e6c38704e4d5896bbb05afb22924c07ae8ae --- /dev/null +++ b/model-00292-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:e77e013cb9fcea93f0423b391aad8937cc7373911c1e795077068405b34dfcfc +size 2822748248 diff --git a/model-00293-of-00341.safetensors b/model-00293-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..6e327f55f96663063712c791df422dd7e4239600 --- /dev/null +++ b/model-00293-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:34620746a4f88691036f2c3627fdd64d48c7ef66f414e43013d2734f5c327b23 +size 3152265592 diff --git a/model-00294-of-00341.safetensors b/model-00294-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..ba65e689b626b5ffd563894aabc612aa82c04723 --- /dev/null +++ b/model-00294-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0d952f15be2fa1572d822e8da98f621d027c218f587396f4e47869e4333f41f7 +size 3153288400 diff --git a/model-00295-of-00341.safetensors b/model-00295-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..ec0660b8cfbb628d18bd31aa30008232e1169a5d --- /dev/null +++ b/model-00295-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fc4f3ef731a1e1fa4b4cb35632e5227ee313f7e5b9bcc520dcc47628663eca02 +size 2401529096 diff --git a/model-00296-of-00341.safetensors b/model-00296-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..5fd0e90dc47f07f7c2e8abba7e652d6030610839 --- /dev/null +++ b/model-00296-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9a9d3cdc331aa47c0b41f10b091e5c89fcb41c38cd9682ee84848697e835cade +size 3154458912 diff --git a/model-00297-of-00341.safetensors b/model-00297-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..eab2cd4bda1fd960050ed4fbc58b1728d1ddbd48 --- /dev/null +++ b/model-00297-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:21a178e8bfe8ef7728917cb1c7264e7b1831d718aa5128a0af93f6559b84582b +size 3153288488 diff --git a/model-00298-of-00341.safetensors b/model-00298-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..156348ce9c6c5e7961a83cce71871982b85dc09a --- /dev/null +++ b/model-00298-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:76a725fee3fd3944bee891cdf6ac4874075e865c77d63d7afa0970d855d01f66 +size 2822748248 diff --git a/model-00299-of-00341.safetensors b/model-00299-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..0fb00f3b93b0b7869363848a2513b812be719365 --- /dev/null +++ b/model-00299-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c5f84c85652deec2fc7b7efd4dca9cba12ace7e7b269af5e046c2df57400238f +size 3154458912 diff --git a/model-00300-of-00341.safetensors b/model-00300-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..bc7825946eeef9646ddb955927e0b43b0f2af016 --- /dev/null +++ b/model-00300-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6e845ffe86d7c02e6abf657d2b580bd047e26eb4cb261ef87cb7e3dc56101b25 +size 3153288488 diff --git a/model-00301-of-00341.safetensors b/model-00301-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c91053b398d0652ad8d24a9deda5c73d7e605063 --- /dev/null +++ b/model-00301-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:98a0e7dfd1cb27039431d94683c80d0c39d00ecff32749cedcbe8d46a40155e4 +size 2822748248 diff --git a/model-00302-of-00341.safetensors b/model-00302-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..be08d4d890dd454c4371f017f54b42b9c4f52e54 --- /dev/null +++ b/model-00302-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:40a12d1220da95eff9c88f8ba9ff53b2f2a168228c7de9ef63f77136efe8ffe9 +size 3154458912 diff --git a/model-00303-of-00341.safetensors b/model-00303-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c115a334a94ce8db0c6bfdace2902bd9ddee69d1 --- /dev/null +++ b/model-00303-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:56e7aeb737c61a093881b3a3fc2e99b66722d6771738533bc9f1cb285dcfcf2b +size 3153288488 diff --git a/model-00304-of-00341.safetensors b/model-00304-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..b78af5e69decdd81ebacf30ee8624437812599c0 --- /dev/null +++ b/model-00304-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:051f3b5007c7437e90009c1f7effbfa82a5f66e596d6d6271095daa1b9e0a89f +size 2822748248 diff --git a/model-00305-of-00341.safetensors b/model-00305-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c8ea7a01dfa0a890ecd568d31622cf6362937c98 --- /dev/null +++ b/model-00305-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:748903342051e7dbf148c614cd624b517227b83d2e12e3955f012d80de6100f6 +size 3152265592 diff --git a/model-00306-of-00341.safetensors b/model-00306-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..ffde88843039c73b9bb61c9f665e15698bceced2 --- /dev/null +++ b/model-00306-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1827e511f01bf5ab322c14c9fba09cd6387d241efa1379597abb3c6941ff7b99 +size 3153288400 diff --git a/model-00307-of-00341.safetensors b/model-00307-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..267cbd897bd7771909ef3ff98715869de04c8258 --- /dev/null +++ b/model-00307-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:dde7f54043bb9b2b2f6b74045a47b4b2c7ed176fbb42ef383b88de37632589d0 +size 2401529096 diff --git a/model-00308-of-00341.safetensors b/model-00308-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..4319524459d756bbaddd97a39432957917a3fbd8 --- /dev/null +++ b/model-00308-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1437936b0678d6184c69c031f176bbc61e8d4b251961f67ad64b81ab3e32f977 +size 3154458912 diff --git a/model-00309-of-00341.safetensors b/model-00309-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..ecf89cedcf4b40765d858fb9fd78e2d9d7283b24 --- /dev/null +++ b/model-00309-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7be7de7469e5e56aec099034ea2204a092cebfb984274a04ae89d42580cf4fb6 +size 3153288488 diff --git a/model-00310-of-00341.safetensors b/model-00310-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..0d688182a8ae233475ae2126b49fc5a4bf6243c4 --- /dev/null +++ b/model-00310-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d3dbbafe6424f689187f5bb51684e023d9f9e7fbdef71c1dd9f2805005689273 +size 2822748248 diff --git a/model-00311-of-00341.safetensors b/model-00311-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..6c5f05940e675ca1dda7fc13057611d95b395129 --- /dev/null +++ b/model-00311-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2f6ed835c8dd04639d08e85a69895b3b0ad773e35a4a3e1039ce52cb25362bdc +size 3154458912 diff --git a/model-00312-of-00341.safetensors b/model-00312-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..d47c4b018960ca152d484ff326a6297d97cd650b --- /dev/null +++ b/model-00312-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f0894ea9d91282d6c9dff592bc8f9a7edc053f698c3a098f95a26a28caefadad +size 3153288488 diff --git a/model-00313-of-00341.safetensors b/model-00313-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c7175457678c9e9cce5d46904d6f3f72da808051 --- /dev/null +++ b/model-00313-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:63679beca6cba94a32c046ba13837826e239a8c0380479231be6de1c46138114 +size 2822748248 diff --git a/model-00314-of-00341.safetensors b/model-00314-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..5cb1d4af22b774c93b8d21864e813402e39af993 --- /dev/null +++ b/model-00314-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b885f6f89fd3f87b17054f64107d4ef8400d36401ddb7d9f84a280477e380a5a +size 3154458912 diff --git a/model-00315-of-00341.safetensors b/model-00315-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..7a7a3d17904f0c24c43ebce5da92b9c0a3218a00 --- /dev/null +++ b/model-00315-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:933e118fbd578a18d7985cf8d21a63ce7fb0d7ee97428bb576eebb3da1c7dcd6 +size 3153288488 diff --git a/model-00316-of-00341.safetensors b/model-00316-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..ab16a2196fea7f2ba8c4e059ca4ef25fbb5a6134 --- /dev/null +++ b/model-00316-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:52c9566a33b975804a418df70e52724411227ff01647f4063ef4dc6782c7b746 +size 2822748248 diff --git a/model-00317-of-00341.safetensors b/model-00317-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..0883d35caac4634aeee3623a7f37e87d5f444e3a --- /dev/null +++ b/model-00317-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4113f0972356f639d0b2992ec78f467623b272bd5543225818cd167e3138eb7b +size 3152265592 diff --git a/model-00318-of-00341.safetensors b/model-00318-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..13c788eab3f197ec68f9c67b017dde36fe5cb519 --- /dev/null +++ b/model-00318-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0730d99d4d1bb159ca56f12e60d136460381f315d6fb612cd6864bdc3667dfd8 +size 3153288400 diff --git a/model-00319-of-00341.safetensors b/model-00319-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..1c1a874785b9ca5245876ba59ab076655d06a216 --- /dev/null +++ b/model-00319-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4a420047b273d3cd5d94e550597e5b0ddc6d9ff33b1c39beb93c4e7a8eddc016 +size 2401529096 diff --git a/model-00320-of-00341.safetensors b/model-00320-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..84b90f2546d6ed84e9bd2da8e1ced942db04c604 --- /dev/null +++ b/model-00320-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d7c6aa8dff09528b1a16d91609993ea3bb4b3b8201b10479bf41258301d2b8f +size 3154458912 diff --git a/model-00321-of-00341.safetensors b/model-00321-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..250f8f5deb15fd282677c9c698921bc527c78cc7 --- /dev/null +++ b/model-00321-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:905686683836c7c8357c3ee423645c891ffa1f6bcbfedac14c33b69905a7d8c5 +size 3153288488 diff --git a/model-00322-of-00341.safetensors b/model-00322-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..7533ef5b0a7de0f6270d5ccd3e317bf2a6978e05 --- /dev/null +++ b/model-00322-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fb299c8ac8b5ad14cacf6204f5e180d47263e8946f1a765de43ad2f2ce21728a +size 2822748248 diff --git a/model-00323-of-00341.safetensors b/model-00323-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..1b61f5183a556ca6158c9413bc4d45eda6584bcf --- /dev/null +++ b/model-00323-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:87e80dfe578666a538727683759e944e822c5af0bcdffa60a1fe7084a3401878 +size 3154458912 diff --git a/model-00324-of-00341.safetensors b/model-00324-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..dd2b477a1c0c53c99aff28224cc21465294b7edd --- /dev/null +++ b/model-00324-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f40f2a482edfe329c4c2904a4f2c053e858fab6078e036ed1929c6e21f796c24 +size 3153288488 diff --git a/model-00325-of-00341.safetensors b/model-00325-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..cab2355ada1650dc54cc188742b2a57507f6613e --- /dev/null +++ b/model-00325-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:6b7320a563f35770461e59c81947fec994bd19f9d9ec52962a54d3dc2ea727f0 +size 2822748248 diff --git a/model-00326-of-00341.safetensors b/model-00326-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..00946bf18993715fae79f0c4750b602fd241f18a --- /dev/null +++ b/model-00326-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:094ffc897f110c6af4233e0aa845e663e4a0e99f51e1793b9281f168f8de506f +size 3154458912 diff --git a/model-00327-of-00341.safetensors b/model-00327-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..4db477e34b8159dee690a739a5a47420c1a2c0dc --- /dev/null +++ b/model-00327-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:073b38b1bad824d2c5f3f0c0deb0504e3080fc6ff2eb00f1effb761a5e53f7a3 +size 3153288488 diff --git a/model-00328-of-00341.safetensors b/model-00328-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..76818bc13b778731fd59a38faa579d05b0661b72 --- /dev/null +++ b/model-00328-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1078837e3dbbcf8ee2fba966a358bcbcfcc24279e479d77b94bf3e9a9a0adccd +size 2822748248 diff --git a/model-00329-of-00341.safetensors b/model-00329-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c8311af59e6a22d3d8a118333e53d5af985155eb --- /dev/null +++ b/model-00329-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:eb8c4a2df02c33256e6c20c37c3ee8cad14b9c5726ea843ec084e9dce2f9e784 +size 3152265592 diff --git a/model-00330-of-00341.safetensors b/model-00330-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..63a36bead7d36f1cc1f7d6af81b67a524b84e236 --- /dev/null +++ b/model-00330-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4c1f765c9303b881f46f54879956c888ffc6f7cd611ba7854dfc1a83ea63265e +size 3153288400 diff --git a/model-00331-of-00341.safetensors b/model-00331-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..ddca8df924f837e458978906fea2a102e13b92ae --- /dev/null +++ b/model-00331-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be66a3b841d1e8b8c0c9c89ae697688d6347b4857db38563e984bf98d8839e4b +size 2401529096 diff --git a/model-00332-of-00341.safetensors b/model-00332-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..682a34d285ddf7f0239ca929eec58da173d8e3e6 --- /dev/null +++ b/model-00332-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:739e13d4f4ab03dc356612687dfa5fc8e58d22934ef09b35c2cebc5ef79d53a8 +size 3149461232 diff --git a/model-00333-of-00341.safetensors b/model-00333-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..96b0336f1673fd36cf7e0b230feb0c377db7a27c --- /dev/null +++ b/model-00333-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:15779f95c4347bfe3d7949a590e4b8a6cc43f8c1b973acdd645aaf8f4d32daae +size 3151088816 diff --git a/model-00334-of-00341.safetensors b/model-00334-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..b9546106b2bc7cc37c9cde2b9ba1a8a090390118 --- /dev/null +++ b/model-00334-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f99f55ad7e25192874841979e01b6a807190ea8976c28a1419fe4a510279e221 +size 3151088832 diff --git a/model-00335-of-00341.safetensors b/model-00335-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..a40f7602c9bec05a91b74579a7c787584040abce --- /dev/null +++ b/model-00335-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4c136cbcfa3a481167f552478c730e8715a898e383583a8668598d4abe3e2ad3 +size 3151088736 diff --git a/model-00336-of-00341.safetensors b/model-00336-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..c98e9c0e60d229568065579e24e2416c1052a3df --- /dev/null +++ b/model-00336-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:96df2654cf46b4586264c1a0a87db113f908ddd58ed598a81f1c0b4740826b02 +size 3151088824 diff --git a/model-00337-of-00341.safetensors b/model-00337-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..a56d1eb3b5812902107481f843b169eaed0f634b --- /dev/null +++ b/model-00337-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8bcd954995ab2d269cfe0f295f462dfe29c5d439a97c58faa136c3e4ad4d2bbd +size 352013032 diff --git a/model-00338-of-00341.safetensors b/model-00338-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..33b3b6fb3b30389c310b53157ae9dcb2fef4dcef --- /dev/null +++ b/model-00338-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:100cb9bf216e214fa40f19172a258298d6253037b50218cb0468a759e0a80b5a +size 2348810352 diff --git a/model-00339-of-00341.safetensors b/model-00339-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..0b9974b9f3ba4ce6fabf828ebae1ede2bb7c75c2 --- /dev/null +++ b/model-00339-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:aa739185387fb8e0d7fe5a19e5cad5249f9b24a028f498f4009d85adf0e553d9 +size 2348853728 diff --git a/model-00340-of-00341.safetensors b/model-00340-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..edb58d17e62992bd1817de893897e9083d82547b --- /dev/null +++ b/model-00340-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:01d41139abb8cf3b5288a97318cc4ab92676671b4eb141031cd80d7c3ced6122 +size 92289328 diff --git a/model-00341-of-00341.safetensors b/model-00341-of-00341.safetensors new file mode 100644 index 0000000000000000000000000000000000000000..da296080d7800b10d6f98c8b5e093cb35923bd23 --- /dev/null +++ b/model-00341-of-00341.safetensors @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9d10c74fc10161bef9463a8541a634a97f521f43c99368ea7243ce0c79cdbf7c +size 802448352 diff --git a/model.safetensors.index.json b/model.safetensors.index.json new file mode 100644 index 0000000000000000000000000000000000000000..c3a3aa6749a2bb609d69ae28d34f1ca06932a3d0 --- /dev/null +++ b/model.safetensors.index.json @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:1ccd22b241272c52b869c8627b6984168396e65e9a299032851661b042e7a6f5 +size 116035852 diff --git a/modeling_kimi_k3.py b/modeling_kimi_k3.py new file mode 100644 index 0000000000000000000000000000000000000000..be85612710534c2e41065b67933a9639e7817d7f --- /dev/null +++ b/modeling_kimi_k3.py @@ -0,0 +1,1317 @@ +# coding=utf-8 +# Copyright 2025-2026 The Moonshot AI Team and HuggingFace Inc. team. All rights reserved. +# +# The code is based on llava (llava/modeling_llava.py), but modified for Kimi-K3. +# +# Licensing Information: +# - Code derived from llava (llava/modeling_llava.py) is licensed under the Apache License, Version 2.0. +# - Other parts of the code are licensed under the Kimi K3 License (see the LICENSE file in this repository). +# +# Apache License, Version 2.0: +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + + +# NOTE: Reference implementation for model architecture; see the model card for production deployment. +import math +from collections.abc import Sequence +from copy import deepcopy +from typing import Optional + +import numpy as np +import torch +import torch.nn as nn +import torch.nn.functional as F +from transformers import activations + +try: + from transformers.activations import PytorchGELUTanh +except ImportError: + from transformers.activations import GELUTanh + activations.PytorchGELUTanh = GELUTanh + PytorchGELUTanh = GELUTanh +from transformers.activations import PytorchGELUTanh +from transformers.configuration_utils import PretrainedConfig +from transformers.modeling_utils import PreTrainedModel +from transformers.models.llava.modeling_llava import \ + LlavaCausalLMOutputWithPast +from transformers.utils import is_flash_attn_2_available + +from .configuration_kimi_k3 import KimiK3Config +from .modeling_kimi_linear import KimiLinearForCausalLM + +# Flash attention imports +if is_flash_attn_2_available(): + from flash_attn import flash_attn_varlen_func +else: + flash_attn_varlen_func = None + + +def multihead_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_cu_seqlens: torch.Tensor | None = None, + k_cu_seqlens: torch.Tensor | None = None, + max_seqlen_q: int | None = None, + max_seqlen_k: int | None = None, + deterministic: bool = False, +): + """Multi-head attention using flash attention 2. + + Args: + q, k, v: tensor of shape (batch_size, seqlen, num_heads, head_dim), + or (tot_seqlens, num_heads, head_dim) if packing. + q_cu_seqlens (torch.Tensor): cumulative sequence lengths of q. + The first element should be 0 and the last element should be q.shape[0]. + k_cu_seqlens (torch.Tensor): cumulative sequence lengths of k. + The first element should be 0 and the last element should be k.shape[0]. + + Returns: + output: shape (batch_size, seqlen, dim) or (tot_seqlens, dim) if packing, + where dim = num_heads * head_dim + """ + attn_out = flash_attn_varlen_func( + q, + k, + v, + q_cu_seqlens, + k_cu_seqlens, + max_seqlen_q, + max_seqlen_k, + causal=False, + deterministic=deterministic, + ) + if isinstance(attn_out, tuple): + attn_out = attn_out[0] + + attn_out = attn_out.flatten(start_dim=-2) + + return attn_out + + +def eager_attention( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + q_cu_seqlens: Optional[torch.Tensor] = None, + k_cu_seqlens: Optional[torch.Tensor] = None, + **kwargs, +) -> torch.Tensor: + seq_length = q.shape[0] + attention_mask = torch.zeros([1, seq_length, seq_length], + device=q.device, + dtype=torch.bool) + for i in range(1, len(q_cu_seqlens)): + attention_mask[ + ..., + q_cu_seqlens[i - 1]:q_cu_seqlens[i], + q_cu_seqlens[i - 1]:q_cu_seqlens[i], + ] = True + q = q.transpose(0, 1) + k = k.transpose(0, 1) + v = v.transpose(0, 1) + + attn_weight = q @ k.transpose(-2, -1) / math.sqrt(q.shape[-1]) + attn_weight = attn_weight.masked_fill( + ~attention_mask, torch.finfo(attn_weight.dtype).min) + attn_weight = torch.softmax(attn_weight, dim=-1, + dtype=torch.float32).to(q.dtype) + + attn_output = attn_weight @ v + attn_output = attn_output.transpose(0, 1) + attn_output = attn_output.reshape(seq_length, -1) + return attn_output + + +VL_VISION_ATTENTION_FUNCTIONS = { + "flash_attention_2": multihead_attention, + "eager": eager_attention, +} + + +def _apply_rope_input_validation(x, freqs_cis): + assert x.ndim == freqs_cis.ndim + 1, (x.shape, freqs_cis.shape) + assert x.shape[:-2] == freqs_cis.shape[:-1], (x.shape, freqs_cis.shape) + assert x.shape[-1] == 2 * freqs_cis.shape[-1], (x.shape, freqs_cis.shape) + assert freqs_cis.dtype == torch.complex64, freqs_cis.dtype + + +def get_rope_shape_decorate(func): + _get_rope_shape_first_call_flag = set() + + def wrapper(org, interpolation_mode, shape): + key = (org.requires_grad, torch.is_grad_enabled(), interpolation_mode) + if key not in _get_rope_shape_first_call_flag: + _get_rope_shape_first_call_flag.add(key) + _ = func(org, interpolation_mode, shape=(64, 64)) + return func(org, interpolation_mode, shape) + + return wrapper + + +@get_rope_shape_decorate +@torch.compile(dynamic=True) +def get_rope_shape(org, interpolation_mode, shape): + return (F.interpolate( + org.permute((2, 0, 1)).unsqueeze(0), + size=shape, + mode=interpolation_mode, + ).squeeze(0).permute((1, 2, 0)).flatten(end_dim=1)) + + +def apply_rope(xq: torch.Tensor, xk: torch.Tensor, + freqs_cis: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: + """ + Args: (The leading dimensions of all inputs should be the same) + xq: query, tensor of shape (..., num_heads, head_dim) + xk: key, tensor of shape (..., num_heads, head_dim) + freqs_cis: tensor of shape (..., head_dim/2), dtype=torch.complex64. It contains the precomputed cis(freqs) for each position in the 2D grid. + Returns: + xq_out, xk_out: tensors of shape (..., num_heads, head_dim) + """ + _apply_rope_input_validation(xq, freqs_cis) + _apply_rope_input_validation(xk, freqs_cis) + + freqs_cis = freqs_cis.unsqueeze(-2) # ..., 1, head_dim/2 + # ..., num_heads, head_dim/2 + xq_ = torch.view_as_complex(xq.float().view(*xq.shape[:-1], -1, 2)) + xk_ = torch.view_as_complex(xk.float().view(*xq.shape[:-1], -1, 2)) + xq_out = torch.view_as_real(xq_ * freqs_cis).flatten( + -2) # ..., num_heads, head_dim + xk_out = torch.view_as_real(xk_ * freqs_cis).flatten( + -2) # ..., num_heads, head_dim + return xq_out.type_as(xq), xk_out.type_as(xk) + + +def get_1d_sincos_pos_embed_from_grid(embed_dim, pos): + """ + From: + https://github.com/OpenGVLab/InternVideo/blob/421f6d2361fc8f61a3394244571f2601a4e99e29/InternVideo2/multi_modality/models/backbones/internvideo2/pos_embed.py#L86 + embed_dim: output dimension for each position + pos: a list of positions to be encoded: size (M,) + out: (M, D) + """ + assert embed_dim % 2 == 0 + omega = np.arange(embed_dim // 2, dtype=np.float32) + omega /= embed_dim / 2.0 + omega = 1.0 / 10000**omega # (D/2,) + + pos = pos.reshape(-1) # (M,) + out = np.einsum('m,d->md', pos, omega) # (M, D/2), outer product + + emb_sin = np.sin(out) # (M, D/2) + emb_cos = np.cos(out) # (M, D/2) + + emb = np.concatenate([emb_sin, emb_cos], axis=1) # (M, D) + return emb + + +def get_1d_sincos_pos_embed(embed_dim, t_size, cls_token=False): + """ + t_size: int of the temporal size + return: + pos_embed: [t_size, embed_dim] or [1+t_size, embed_dim] (w/ or w/o cls_token) + """ + grid_t = np.arange(t_size, dtype=np.float32) + pos_embed = get_1d_sincos_pos_embed_from_grid(embed_dim, grid_t) + if cls_token: + pos_embed = np.concatenate([np.zeros([1, embed_dim]), pos_embed], + axis=0) + return pos_embed + + +class Learnable2DInterpPosEmbDivided_fixed(nn.Module): + + def __init__(self, + height: int, + width: int, + num_frames: int, + dim: int, + interpolation_mode: str = 'bicubic') -> None: + super().__init__() + self.height = height + self.width = width + self.num_frames = num_frames + self.dim = dim + self.interpolation_mode = interpolation_mode + self.weight = nn.Parameter(torch.empty(height, width, dim)) + self.register_buffer('time_weight', + torch.from_numpy( + get_1d_sincos_pos_embed( + self.dim, + self.num_frames)).float().unsqueeze(1), + persistent=False) + + self.reset_parameters() + + def reset_parameters(self): + nn.init.normal_(self.weight) + + def forward(self, x: torch.Tensor, + grid_thws: torch.Tensor) -> torch.Tensor: + pos_embs = [] + for t, h, w in grid_thws.tolist(): + assert t <= self.num_frames, f't:{t} > self.num_frames:{self.num_frames}' + if (h, w) == self.weight.shape[:-1]: + pos_emb_2d = self.weight.flatten(end_dim=1) + else: + pos_emb_2d = get_rope_shape( + self.weight, + interpolation_mode=self.interpolation_mode, + shape=(h, w), + ) + + if t == 1: + pos_emb_3d = pos_emb_2d + else: + pos_emb_3d = pos_emb_2d.unsqueeze(0).repeat( + t, 1, 1) + self.time_weight[0:t] + + pos_embs.append(pos_emb_3d.reshape(-1, pos_emb_3d.shape[-1])) + + out = x + torch.cat(pos_embs) + return out + + +class MoonVision3dPatchEmbed(nn.Module): + + def __init__(self, + out_dim: int, + in_dim: int = 3, + patch_size: int | tuple[int, int] = (14, 14), + pos_emb_height: int = 14, + pos_emb_width: int = 14, + pos_emb_time: int = 4, + pos_emb_type: str = 'divided_fixed', + patch_embed_proj_bias: bool = True, + pos_emb_interpolation_mode: str = 'bicubic'): + super().__init__() + assert isinstance( + patch_size, + int | Sequence), f'Invalid patch_size type: {type(patch_size)}' + if isinstance(patch_size, int): + patch_size = (patch_size, patch_size) + assert (len(patch_size) == 2 + ), f'Expected patch_size to be a tuple of 2, got {patch_size}' + self.patch_size = patch_size + + self.proj = nn.Conv2d(in_dim, + out_dim, + kernel_size=patch_size, + stride=patch_size, + bias=patch_embed_proj_bias) + + if pos_emb_type == 'divided_fixed': + self.pos_emb = Learnable2DInterpPosEmbDivided_fixed( + height=pos_emb_height, + width=pos_emb_width, + num_frames=pos_emb_time, + dim=out_dim, + interpolation_mode=pos_emb_interpolation_mode) + else: + raise NotImplementedError( + f'Not support pos_emb_type: {pos_emb_type}') + + def forward(self, x: torch.Tensor, + grid_thws: torch.Tensor) -> torch.Tensor: + """ + Args: + x (L, Channels): input tensor + grid_hws (N, 3): temporal, height and width + + Returns: + (L, Cout) tensor + """ + x = self.proj(x).view(x.size(0), -1) + # apply positional embedding + x = self.pos_emb(x, grid_thws) + return x + + +class Rope2DPosEmbRepeated(nn.Module): + """2D rotary position embedding with multi-resolution support. + + This class is intended to be used in the following way: + 1. Before training, create an instance of Rope2DPosEmb. This instance will hold the precomputed cis. + 2. Before each forward pass, call `get_freqs_cis_by_*` to get the `freqs_cis` tensor for this iteration. + 3. During the forward pass, pass the `freqs_cis` tensor to each attention layer, and call `apply` just before each attention operation. + The rope is shared across all attention layers and all heads. + + Refs: + - RoFormer: https://arxiv.org/abs/2104.09864 + - VisionLLaMA: https://arxiv.org/abs/2403.00522 + - https://github.com/Meituan-AutoML/VisionLLaMA/blob/main/dit/models.py + + Args: + dim (int): usually the multi-head attention dimension, should be divisible by 4 (TODO: relax this constraint if needed) + max_height (int): the maximum height of the 2D grid + max_width (int): the maximum width of the 2D grid + theta_base (float): the base of the theta + device (str): the device to store the precomputed cis + """ + + def __init__(self, + dim: int, + max_height: int, + max_width: int, + theta_base=10000): + super().__init__() + self.dim = dim + assert self.dim % 4 == 0, 'dim must be divisible by 4' + self.max_height = max_height + self.max_width = max_width + self.theta_base = theta_base + + def extra_repr(self): + return f'dim={self.dim}, max_height={self.max_height}, max_width={self.max_width}, theta_base={self.theta_base}' + + def _precompute_freqs_cis(self, device: torch.device) -> torch.Tensor: + """Calculate the cis(freqs) for each position in the 2D grid. + + Return: complex tensor of shape (max_height, max_width, dim//2) and value: + height axis: ret[h, w, 2*i] = cis(h * theta_base**(-4*i/dim)) + weight axis: ret[h, w, 2*i+1] = cis(w * theta_base**(-4*i/dim)) with (i in [0, dim//4)) + note: `cis` is a mathematical notation defined by cis x = cos x + i sin x, + """ + N = self.max_height * self.max_width + flat_pos = torch.arange(0, N).float().to(device) + x_pos = flat_pos % self.max_width + y_pos = flat_pos // self.max_width + dim_range = (torch.arange(0, self.dim, + 4)[:(self.dim // 4)].float().to(device) + ) # C/4 + freqs = 1.0 / (self.theta_base**(dim_range / self.dim)) + x_freqs = torch.outer(x_pos, freqs).float() # N, C/4 + y_freqs = torch.outer(y_pos, freqs).float() # N, C/4 + x_cis = torch.polar(torch.ones_like(x_freqs), x_freqs) # N, C/4 + y_cis = torch.polar(torch.ones_like(y_freqs), y_freqs) # N, C/4 + # N, C/4, 2 + freqs_cis = torch.cat( + [x_cis.unsqueeze(dim=-1), + y_cis.unsqueeze(dim=-1)], dim=-1) + # max_height, max_width, C/2 + freqs_cis = freqs_cis.reshape(self.max_height, self.max_width, -1) + return freqs_cis + + def get_freqs_cis(self, grid_thws: torch.Tensor, + device: torch.device) -> torch.Tensor: + """ + Args: + grid_thws (torch.Tensor): grid time, height and width + + Returns: + freqs_cis: tensor of shape (sum(t * height * width), dim//2) + """ + if not hasattr(self, 'freqs_cis'): + self.register_buffer('freqs_cis', + self._precompute_freqs_cis(device), + persistent=False) + + shapes = grid_thws.tolist() + assert all(1 <= h <= self.max_height and 1 <= w <= self.max_width + for t, h, w in shapes), ( + shapes, + self.max_height, + self.max_width, + ) + freqs_cis = torch.cat( + [ + self.freqs_cis[:h, :w].reshape(-1, self.dim // 2).repeat(t, 1) + for t, h, w in shapes + ], + dim=0, + ) + return freqs_cis + + +class MLP2(nn.Module): + """ + Args: + dims: [in_dim, hidden_dim, out_dim] + bias: whether to use bias in linear layer. + """ + + def __init__(self, dims: list[int], activation, bias=True): + super().__init__() + assert len(dims) == 3 + self.fc0 = nn.Linear(dims[0], dims[1], bias=bias) + self.fc1 = nn.Linear(dims[1], dims[2], bias=bias) + self.activation = activation + for m in [self.fc0, self.fc1]: + nn.init.trunc_normal_(m.weight, std=math.sqrt(2 / m.in_features)) + if m.bias is not None: + nn.init.zeros_(m.bias) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + x = self.fc0(x) + x = self.activation(x) + return self.fc1(x) + + +class MoonViTEncoderLayer(nn.Module): + + def __init__( + self, + num_heads: int, + hidden_dim: int, + mlp_dim: int, + qkv_hidden_size: int | None = None, + norm_type: str = 'layernorm', + mlp_type: str = 'mlp2', + *, + attn_implementation: str = 'flash_attention_2', + activation=F.gelu, + attn_bias: bool = False, + linear_bias: bool = True, + use_deterministic_attn: bool = False, + ): + super().__init__() + self.num_heads = num_heads + self.hidden_dim = hidden_dim + self.qkv_hidden_size = hidden_dim if qkv_hidden_size is None else qkv_hidden_size + self.hidden_size_per_attention_head = self.qkv_hidden_size // self.num_heads + self.attn_implementation = attn_implementation + self.use_deterministic_attn = use_deterministic_attn + + if norm_type == "layernorm": + self.norm0 = nn.LayerNorm(hidden_dim) + self.norm1 = nn.LayerNorm(hidden_dim) + elif norm_type == "rmsnorm": + self.norm0 = nn.RMSNorm(hidden_dim) + self.norm1 = nn.RMSNorm(hidden_dim) + else: + raise NotImplementedError(f"Not support norm_type: {norm_type}") + + if mlp_type == "mlp2": + self.mlp = MLP2([hidden_dim, mlp_dim, hidden_dim], + activation, + bias=linear_bias) + else: + raise NotImplementedError(f"Not support mlp_type: {mlp_type}") + + self.wqkv = nn.Linear(hidden_dim, + self.qkv_hidden_size * 3, + bias=attn_bias) + self.wo = nn.Linear(self.qkv_hidden_size, hidden_dim, bias=attn_bias) + + def attention_qkvpacked( + self, + x: torch.Tensor, + cu_seqlens: torch.Tensor, + max_seqlen: torch.Tensor, + rope_freqs_cis: torch.Tensor | None = None, + ): + """ + Args: + x (torch.Tensor): (batch_size, seqlen, hidden_dim) + cu_seqlens (torch.Tensor): + """ + xqkv = self.wqkv(x) + + qkv_shape = xqkv.size()[:-1] + ( + 3, + self.num_heads, + self.hidden_size_per_attention_head, + ) + # xqkv: (batch_size, seqlen, 3, nheads, headdim) + xqkv = xqkv.view(*qkv_shape) + xq, xk, xv = torch.unbind(xqkv, dim=-3) + + xq, xk = apply_rope(xq, xk, rope_freqs_cis) + + attn_func = VL_VISION_ATTENTION_FUNCTIONS[self.attn_implementation] + attn_out = attn_func(xq, + xk, + xv, + q_cu_seqlens=cu_seqlens, + k_cu_seqlens=cu_seqlens, + max_seqlen_k=max_seqlen, + max_seqlen_q=max_seqlen, + deterministic=self.use_deterministic_attn) + + attn_out = self.wo(attn_out) + return attn_out + + def forward( + self, + hidden_states: torch.Tensor, + cu_seqlens: torch.Tensor, + max_seqlen: int, + rope_freqs_cis: torch.Tensor | None = None, + ): + residual = hidden_states + hidden_states = self.norm0(hidden_states) + + hidden_states = self.attention_qkvpacked(hidden_states, cu_seqlens, + max_seqlen, rope_freqs_cis) + hidden_states = residual + hidden_states + + residual = hidden_states + hidden_states = self.norm1(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + + return hidden_states + + +class MoonViT3dEncoder(nn.Module): + + def __init__(self, + hidden_dim: int, + num_layers: int, + block_cfg: dict, + use_deterministic_attn: bool = False) -> None: + super().__init__() + self.use_deterministic_attn = use_deterministic_attn + + qkv_hidden_size = block_cfg['hidden_dim'] if block_cfg.get( + 'qkv_hidden_size') is None else block_cfg['qkv_hidden_size'] + self.rope_2d = Rope2DPosEmbRepeated( + qkv_hidden_size // block_cfg['num_heads'], 512, 512) + self.blocks = nn.ModuleList([ + MoonViTEncoderLayer( + **block_cfg, + use_deterministic_attn=self.use_deterministic_attn) + for _ in range(num_layers) + ]) + norm_type = block_cfg.get('norm_type', 'layernorm') + if norm_type == "layernorm": + self.final_layernorm = nn.LayerNorm(hidden_dim) + elif norm_type == "rmsnorm": + self.final_layernorm = nn.RMSNorm(hidden_dim) + else: + raise NotImplementedError(f"Not support norm_type: {norm_type}") + + def forward( + self, + hidden_states: torch.Tensor, + grid_thws: torch.Tensor, + ) -> torch.Tensor: + rope_freqs_cis = self.rope_2d.get_freqs_cis( + grid_thws=grid_thws, device=hidden_states.device) + + lengths = torch.cat(( + torch.zeros(1, dtype=grid_thws.dtype, device=grid_thws.device), + grid_thws[:, 0] * grid_thws[:, 1] * grid_thws[:, 2], + )) + + max_seqlen = lengths.max() + cu_seqlens = lengths.to(hidden_states.device).cumsum(dim=0, + dtype=torch.int32) + for block in self.blocks: + hidden_states = block(hidden_states, + cu_seqlens, + max_seqlen, + rope_freqs_cis=rope_freqs_cis) + + hidden_states = self.final_layernorm(hidden_states) + return hidden_states + + +def tpool_patch_merger( + x: torch.Tensor, + grid_thws: torch.Tensor, + merge_kernel_size: tuple[int, int] = (2, 2), +) -> list[torch.Tensor]: + d_model = x.size(-1) + + outputs = [] + pre_sum = 0 + for t, h, w in grid_thws.tolist(): + # Get the current sequence + seq = x[pre_sum:pre_sum + t * h * w] + # Reshape along self.merge_kernel_size and concat to the last dimension + kernel_height, kernel_width = merge_kernel_size + new_height, new_width = h // kernel_height, w // kernel_width + reshaped_seq = seq.view(t, new_height, kernel_height, new_width, + kernel_width, d_model) + reshaped_seq = reshaped_seq.permute(0, 1, + 3, 2, 4, 5).contiguous().mean( + dim=0) # temporal pooling + padded_seq = reshaped_seq.view(new_height * new_width, + kernel_height * kernel_width, -1) + outputs.append(padded_seq) + pre_sum += t * h * w + + return outputs + + +class MoonViT3dPretrainedModel(PreTrainedModel): + config_class = None + model_type = 'moonvit3d' + _no_split_modules = ['MoonViTEncoderLayer'] + _supports_flash_attn_2 = True + _supports_sdpa = True + + def __init__(self, config, *inputs, **kwargs): + super().__init__(config, *inputs, **kwargs) + config = deepcopy(config) + self.merge_kernel_size = config.merge_kernel_size + self.patch_size = config.patch_size + self.merge_type = config.merge_type + + self.patch_embed = MoonVision3dPatchEmbed( + out_dim=config.hidden_size, + patch_size=config.patch_size, + pos_emb_height=config.init_pos_emb_height, + pos_emb_width=config.init_pos_emb_width, + pos_emb_time=config.init_pos_emb_time, + pos_emb_type=config.pos_emb_type, + patch_embed_proj_bias=getattr(config, 'patch_embed_proj_bias', + True), + pos_emb_interpolation_mode=getattr( + config, 'pos_emb_interpolation_mode', 'bicubic'), + ) + + self.encoder = MoonViT3dEncoder( + hidden_dim=config.hidden_size, + num_layers=config.num_hidden_layers, + block_cfg={ + 'num_heads': config.num_attention_heads, + 'hidden_dim': config.hidden_size, + 'qkv_hidden_size': getattr(config, 'qkv_hidden_size', None), + 'mlp_dim': config.intermediate_size, + 'norm_type': getattr(config, 'norm_type', 'layernorm'), + 'mlp_type': getattr(config, 'mlp_type', 'mlp2'), + 'activation': PytorchGELUTanh(), + 'attn_bias': getattr(config, 'attn_bias', True), + 'linear_bias': getattr(config, 'linear_bias', True), + 'attn_implementation': config._attn_implementation, + }, + use_deterministic_attn=getattr(self, 'use_deterministic_attn', + False)) + + def forward(self, pixel_values: torch.Tensor, + grid_thws: torch.Tensor) -> torch.Tensor: + """ + Args: + pixel_values (torch.Tensor): The input pixel values. + grid_thws (torch.Tensor): Temporal, height and width. + + Returns: + torch.Tensor: The output tokens. + """ + # grid_thws = grid_thws.to('cpu') + assert grid_thws.ndim == 2, f'grid_thws should be 2D, got {grid_thws.ndim}' + assert grid_thws.size(1) == 3, f'No support for thw: {grid_thws}' + hidden_states = self.patch_embed(pixel_values, grid_thws) + hidden_states = self.encoder(hidden_states, grid_thws) + if self.merge_type == 'sd2_tpool': # spatial downsampling 2x with temporal pooling all + hidden_states = tpool_patch_merger( + hidden_states, + grid_thws, + merge_kernel_size=self.merge_kernel_size) + else: + raise NotImplementedError(f'Not support {self.merge_type}') + + return hidden_states + + +# ============================================================================ +# MM Projector Helper Classes (from mm_projector/modeling_mm_projectors.py) +# ============================================================================ + + +class IdentityMap(nn.Module): + + def __init__(self): + super().__init__() + + def forward(self, x, *args, **kwargs): + return x + + +class MLP(nn.Module): + + def __init__(self, config): + super().__init__() + # TODO, use faster LayerNorm + self.pre_norm = nn.LayerNorm(config.mm_hidden_size) + self.proj = nn.Sequential( + nn.Linear(config.mm_hidden_size, config.hidden_size), nn.GELU(), + nn.Linear(config.hidden_size, config.hidden_size)) + + def forward(self, x, *args, **kwargs): + assert isinstance(x, + list | tuple), f'x is not a list or tuple: {type(x)}' + lengths = [item.shape[0] for item in x] + x = torch.cat(x, dim=0) + x = self.pre_norm(x) + x = self.proj(x) + x = torch.split(x, lengths, dim=0) + + return x + + +class PatchMergerMLP(nn.Module): + + def __init__(self, config): + super().__init__() + eps = config.projector_ln_eps + self.hidden_size = config.mm_hidden_size * ( + config.merge_kernel_size[0] * config.merge_kernel_size[1]) + self.pre_norm = nn.LayerNorm(config.mm_hidden_size, eps=eps) + self.proj = nn.Sequential( + nn.Linear(self.hidden_size, self.hidden_size), + nn.GELU(), + nn.Linear(self.hidden_size, config.hidden_size), + ) + + def forward(self, x, *args, **kwargs): + if isinstance(x, list) or isinstance(x, tuple): + x = [ + self.proj(self.pre_norm(item).view(item.shape[0], -1)) + for item in x + ] + else: + # B, N, N_k, C = x.shape + B = x.shape[0] + x = self.proj(self.pre_norm(x).view(B, -1, self.hidden_size)) + return x + + +class PatchMergerMLPV2(nn.Module): + + def __init__(self, config): + super().__init__() + eps = config.projector_ln_eps + self.hidden_size = config.mm_hidden_size * ( + config.merge_kernel_size[0] * config.merge_kernel_size[1]) + self.proj = nn.Sequential( + nn.Linear(self.hidden_size, self.hidden_size, bias=False), + nn.GELU(), + nn.Linear(self.hidden_size, config.hidden_size, bias=False), + ) + self.post_norm = nn.RMSNorm(config.hidden_size, eps=eps) + for m in self.proj.modules(): + if isinstance(m, nn.Linear): + nn.init.trunc_normal_(m.weight, + std=math.sqrt(2 / m.in_features)) + if m.bias is not None: + nn.init.zeros_(m.bias) + + def forward(self, x, *args, **kwargs): + if isinstance(x, list) or isinstance(x, tuple): + lengths = [item.shape[0] for item in x] + x = torch.concat([item.view(item.shape[0], -1) for item in x], + dim=0) + x = self.post_norm(self.proj(x)) + x = torch.split(x, lengths, dim=0) + else: + # B, N, N_k, C = x.shape + B = x.shape[0] + x = self.proj(x.view(B, -1, self.hidden_size)) + x = self.post_norm(x) + return x + + +class KimiK3PreTrainedModel(PreTrainedModel): + config_class = KimiK3Config + base_model_prefix = "model" + _no_split_modules = [ + "MoonViT3dPretrainedModel", + "MoonViTEncoderLayer", + "KimiDecoderLayer", + "PatchMergerMLP", + "PatchMergerMLPV2", + ] + _skip_keys_device_placement = "past_key_values" + _supports_flash_attn_2 = True + _supports_sdpa = False + + def _init_weights(self, module): + # important: this ported version of Llava isn't meant for training from scratch - only + # inference and fine-tuning - so the proper init weights code has been removed - the original codebase + # https://github.com/haotian-liu/LLaVA/tree/main/llava should serve for that purpose + std = (self.config.initializer_range if hasattr( + self.config, "initializer_range") else + self.config.text_config.initializer_range) + + if hasattr(module, "class_embedding"): + module.class_embedding.data.normal_(mean=0.0, std=std) + + if isinstance(module, (nn.Linear, nn.Conv2d)): + module.weight.data.normal_(mean=0.0, std=std) + if module.bias is not None: + module.bias.data.zero_() + elif isinstance(module, nn.Embedding): + module.weight.data.normal_(mean=0.0, std=std) + if module.padding_idx is not None: + module.weight.data[module.padding_idx].zero_() + + +class VisionTowerConfig(PretrainedConfig): + model_type = 'moonvit3d' + + def __init__(self, config: KimiK3Config, **kwargs): + super().__init__(**kwargs) + self.patch_size = config.patch_size + self.init_pos_emb_height = config.init_pos_emb_height + self.init_pos_emb_width = config.init_pos_emb_width + self.init_pos_emb_time = config.init_pos_emb_time + self.pos_emb_type = config.pos_emb_type + self.num_attention_heads = config.vt_num_attention_heads + self.num_hidden_layers = config.vt_num_hidden_layers + self.hidden_size = config.vt_hidden_size + self.intermediate_size = config.vt_intermediate_size + self.merge_kernel_size = config.merge_kernel_size + self.merge_type = config.merge_type + self._attn_implementation = config._attn_implementation + self.qkv_hidden_size = getattr(config, 'qkv_hidden_size', None) + self.norm_type = getattr(config, 'norm_type', 'layernorm') + self.attn_bias = getattr(config, 'attn_bias', True) + self.patch_embed_proj_bias = getattr(config, 'patch_embed_proj_bias', + True) + self.mlp_type = getattr(config, 'mlp_type', 'mlp2') + self.linear_bias = getattr(config, 'linear_bias', True) + self.pos_emb_interpolation_mode = getattr( + config, 'pos_emb_interpolation_mode', 'bilinear') + + +class ProjectorConfig: + + def __init__(self, config: KimiK3Config): + self.mm_projector_type = config.mm_projector_type + self.mm_hidden_size = config.mm_hidden_size + self.hidden_size = config.text_hidden_size + self.merge_kernel_size = config.merge_kernel_size + self.projector_hidden_act = config.projector_hidden_act + self.projector_ln_eps = config.projector_ln_eps + + +# ref https://github.com/huggingface/transformers/blob/78b2929c0554b79e0489b451ce4ece14d265ead2/src/transformers/models/llava/modeling_llava.py#L240 +class KimiK3ForConditionalGeneration(KimiK3PreTrainedModel): + + @classmethod + def _supports_default_dynamic_cache(cls) -> bool: + return False + + def __init__(self, config: KimiK3Config): + super().__init__(config) + + vt_config = VisionTowerConfig(config.vision_config) + self.vision_tower = MoonViT3dPretrainedModel(vt_config) + + proj_config = ProjectorConfig(config.vision_config) + if proj_config.mm_projector_type == 'identity': + self.mm_projector = IdentityMap() + elif proj_config.mm_projector_type == 'mlp': + self.mm_projector = MLP(proj_config) + elif proj_config.mm_projector_type == 'patchmerger': + self.mm_projector = PatchMergerMLP(proj_config) + elif proj_config.mm_projector_type == 'patchmergerv2': + self.mm_projector = PatchMergerMLPV2(proj_config) + else: + raise ValueError( + f"Unsupported mm_projector_type: {proj_config.mm_projector_type}" + ) + + self.language_model = KimiLinearForCausalLM(config.text_config) + self.post_init() + + if hasattr(self.language_model, 'dtype'): + target_dtype = self.language_model.dtype + self.vision_tower = self.vision_tower.to(dtype=target_dtype) + self.mm_projector = self.mm_projector.to(dtype=target_dtype) + + def get_input_embeddings(self): + return self.language_model.get_input_embeddings() + + def set_input_embeddings(self, value): + self.language_model.set_input_embeddings(value) + + def get_output_embeddings(self): + return self.language_model.get_output_embeddings() + + def set_output_embeddings(self, new_embeddings): + self.language_model.set_output_embeddings(new_embeddings) + + def set_decoder(self, decoder): + self.language_model.set_decoder(decoder) + + def get_decoder(self): + return self.language_model.get_decoder() + + def tie_weights(self): + return self.language_model.tie_weights() + + def resize_token_embeddings(self, + new_num_tokens: int | None = None, + pad_to_multiple_of=None) -> nn.Embedding: + model_embeds = self.language_model.resize_token_embeddings( + new_num_tokens, pad_to_multiple_of) + # update vocab size + self.config.text_config.vocab_size = model_embeds.num_embeddings + self.vocab_size = model_embeds.num_embeddings + return model_embeds + + def _merge_input_ids_with_image_features( + self, + image_features: list[torch.Tensor], + inputs_embeds: torch.Tensor, + input_ids: torch.Tensor, + attention_mask: torch.Tensor, + labels: torch.Tensor | None = None, + ): + """ + Args: + image_features (:obj:`torch.Tensor` of shape :obj:`(num_image_tokens, embed_dim)`): + The image features to merge with the input embeddings. + inputs_embeds (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length, embed_dim)`): + The input embeddings. + input_ids (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`): + The input ids. + attention_mask (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`): + The attention mask. + labels (:obj:`torch.Tensor` of shape :obj:`(batch_size, sequence_length)`, *optional*): + The labels. + """ + _, embed_dim = image_features[0].shape + feature_lengths = [x.shape[0] for x in image_features] + image_features = torch.cat(image_features, dim=0) + + image_token_index: int = self.config.media_placeholder_token_id + pad_token_id: int = self.config.pad_token_id + ignore_index: int = self.config.ignore_index + + batch_size, sequence_length = input_ids.shape + left_padding = not torch.sum( + input_ids[:, -1] == torch.tensor(pad_token_id)) + + # 1. Create a mask to know where special image tokens are + _token_occupation_table = torch.ones_like(input_ids.flatten()) + _token_occupation_table[input_ids.flatten() == + image_token_index] = torch.tensor( + feature_lengths, + dtype=torch.long, + device=input_ids.device) + _token_occupation_table = _token_occupation_table.reshape( + input_ids.shape) + + max_embed_dim = _token_occupation_table.sum(-1).max().item() + assert ( + max_embed_dim >= sequence_length + ), f"The maximum embedding dimension ({max_embed_dim}) is less than the sequence length ({sequence_length})" + batch_indices, non_image_indices = torch.where( + input_ids != image_token_index) + + # 2. Compute the positions where text should be written + # Calculate new positions for text tokens in merged image-text sequence. + new_token_positions = torch.cumsum(_token_occupation_table, -1) - 1 + nb_image_pad = max_embed_dim - 1 - new_token_positions[:, -1] + if left_padding: + new_token_positions += nb_image_pad[:, + None] # offset for left padding + text_to_overwrite = new_token_positions[batch_indices, + non_image_indices] + + # 3. Create the full embedding, already padded to the maximum position + final_embedding = torch.zeros( + batch_size, + max_embed_dim, + embed_dim, + dtype=inputs_embeds.dtype, + device=inputs_embeds.device, + ) + final_attention_mask = torch.zeros(batch_size, + max_embed_dim, + dtype=attention_mask.dtype, + device=inputs_embeds.device) + if labels is not None: + final_labels = torch.full( + (batch_size, max_embed_dim), + ignore_index, + dtype=input_ids.dtype, + device=input_ids.device, + ) + # In case the Vision model or the Language model has been offloaded to CPU, we need to manually + # set the corresponding tensors into their correct target device. + target_device = inputs_embeds.device + batch_indices, non_image_indices, text_to_overwrite = ( + batch_indices.to(target_device), + non_image_indices.to(target_device), + text_to_overwrite.to(target_device), + ) + attention_mask = attention_mask.to(target_device) + + # 4. Fill the embeddings based on the mask. + final_embedding[batch_indices, + text_to_overwrite] = inputs_embeds[batch_indices, + non_image_indices] + final_attention_mask[batch_indices, + text_to_overwrite] = attention_mask[ + batch_indices, non_image_indices] + if labels is not None: + final_labels[batch_indices, + text_to_overwrite] = labels[batch_indices, + non_image_indices] + + # 5. Fill the embeddings corresponding to the images. Anything that is not `text_positions` needs filling (#29835) + image_to_overwrite = torch.full((batch_size, max_embed_dim), + True, + dtype=torch.bool, + device=inputs_embeds.device) + image_to_overwrite[batch_indices, text_to_overwrite] = False + image_to_overwrite &= image_to_overwrite.cumsum( + -1) - 1 >= nb_image_pad[:, None].to(target_device) + + if image_to_overwrite.sum() != image_features.shape[:-1].numel(): + raise ValueError( + f"The input provided to the model are wrong. The number of image tokens is {image_to_overwrite.sum()} while" + f" the number of image features given to the model is {image_features.shape[:-1].numel()}. " + "This prevents correct indexing and breaks batch generation.") + + final_embedding[image_to_overwrite] = ( + image_features.contiguous().reshape(-1, + embed_dim).to(target_device)) + final_attention_mask |= image_to_overwrite + position_ids = (final_attention_mask.cumsum(-1) - 1).masked_fill_( + (final_attention_mask == 0), 1) + + # 6. Mask out the embedding at padding positions, as we later use the past_key_value value to determine the non-attended tokens. + batch_indices, pad_indices = torch.where(input_ids == pad_token_id) + indices_to_mask = new_token_positions[batch_indices, pad_indices] + + final_embedding[batch_indices, indices_to_mask] = 0 + + if labels is None: + final_labels = None + + return final_embedding, final_attention_mask, final_labels, position_ids + + def _extract_image_features(self, pixel_values: torch.Tensor, + grid_thws: torch.Tensor) -> list[torch.Tensor]: + """ + Args: + pixel_values (:obj:`torch.FloatTensor` of shape :obj:`(batch_size, num_channels, height, width)`): + The pixel values of the images processed by image processor. + grid_thws (:obj:`torch.Tensor` of shape :obj:`(batch_size, 3)`): + The grid, height, width of the images. + + Returns: + selected_image_feature (:obj:`torch.FloatTensor` of shape :obj:`(num_image_tokens, embed_dim)`): + The selected image features to use as input to the projector head. + + """ + + target_dtype = self.vision_tower.patch_embed.proj.weight.dtype + pixel_values = pixel_values.to(target_dtype) + + image_features = self.vision_tower(pixel_values, grid_thws) + return image_features + + def forward( + self, + input_ids: torch.LongTensor | None = None, + pixel_values: torch.FloatTensor | list[torch.FloatTensor] + | None = None, + grid_thws: torch.Tensor | None = None, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: list[torch.FloatTensor] | None = None, + inputs_embeds: torch.FloatTensor | None = None, + labels: torch.LongTensor | None = None, + use_cache: bool | None = None, + output_attentions: bool | None = None, + output_hidden_states: bool | None = None, + return_dict: bool | None = None, + ) -> tuple | LlavaCausalLMOutputWithPast: + r""" + Args: + labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*): + Labels for computing the masked language modeling loss. Indices should either be in `[0, ..., + config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored + (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`. + + ```""" + assert self.vision_tower is not None, "vision_tower is not loaded" + output_attentions = (output_attentions if output_attentions is not None + else self.config.output_attentions) + output_hidden_states = (output_hidden_states + if output_hidden_states is not None else + self.config.output_hidden_states) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + if inputs_embeds is None: + # 1. Extra the input embeddings + inputs_embeds = self.get_input_embeddings()(input_ids) + + # 2. Merge text and images + if pixel_values is not None and len( + pixel_values) > 0 and input_ids.shape[1] != 1: + image_features = self._extract_image_features( + pixel_values, grid_thws) + if self.mm_projector: + image_features = self.mm_projector(image_features) + + inputs_embeds = inputs_embeds.to( + image_features[0].dtype) # num_tokens, embed_dim + inputs_embeds, attention_mask, labels, position_ids = ( + self._merge_input_ids_with_image_features( + image_features, + inputs_embeds, + input_ids, + attention_mask, + labels, + )) + + # In case input_ids.shape[1] == 1 & pixel_values==None & past_key_values != None, we are in the case of + # generation with cache + elif (past_key_values is not None and pixel_values is not None + and input_ids.shape[1] == 1): + # Retrieve the first layer to inspect the logits and mask out the hidden states + # that are set to 0 + first_layer_past_key_value = past_key_values[0][0][:, :, :, 0] + + # Sum all dimensions of head_dim (-2) to avoid random errors such as: https://github.com/huggingface/transformers/pull/28032#issuecomment-1863691941 + batch_index, non_attended_tokens = torch.where( + first_layer_past_key_value.float().sum(-2) == 0) + + # Get the target length + target_length = input_ids.shape[1] + past_length = first_layer_past_key_value.shape[-1] + + extended_attention_mask = torch.ones( + (attention_mask.shape[0], past_length), + dtype=attention_mask.dtype, + device=attention_mask.device, + ) + + # Filter out only the tokens that can be un-attended, this can happen + # if one uses Llava + Fused modules where the cache on the + # first iteration is already big enough, or if one passes custom cache + valid_indices = non_attended_tokens < extended_attention_mask.size( + -1) + new_batch_index = batch_index[valid_indices] + new_non_attended_tokens = non_attended_tokens[valid_indices] + + # Zero-out the places where we don't need to attend + extended_attention_mask[new_batch_index, + new_non_attended_tokens] = 0 + + attention_mask = torch.cat( + (extended_attention_mask, attention_mask[:, + -target_length:]), + dim=1) + position_ids = torch.sum(attention_mask, + dim=1).unsqueeze(-1) - 1 + + outputs = self.language_model( + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + inputs_embeds=inputs_embeds, + use_cache=use_cache, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + ) + + logits = outputs[0] + + loss = None + if labels is not None: + # Shift so that tokens < n predict n + if attention_mask is not None: + shift_attention_mask = attention_mask[..., 1:] + shift_logits = logits[..., :-1, :][shift_attention_mask.to( + logits.device) != 0].contiguous() + shift_labels = labels[..., 1:][shift_attention_mask.to( + labels.device) != 0].contiguous() + else: + shift_logits = logits[..., :-1, :].contiguous() + shift_labels = labels[..., 1:].contiguous() + # Flatten the tokens + loss_fct = nn.CrossEntropyLoss() + loss = loss_fct( + shift_logits.view(-1, shift_logits.size(-1)), + shift_labels.view(-1).to(shift_logits.device), + ) + + if not return_dict: + output = (logits, ) + outputs[1:] + return (loss, ) + output if loss is not None else output + + return LlavaCausalLMOutputWithPast( + loss=loss, + logits=logits, + past_key_values=outputs.past_key_values, + hidden_states=outputs.hidden_states, + attentions=outputs.attentions, + ) + + def prepare_inputs_for_generation( + self, + input_ids, + past_key_values=None, + inputs_embeds=None, + pixel_values=None, + grid_thws=None, + attention_mask=None, + **kwargs, + ): + if past_key_values is not None: + if hasattr(past_key_values, "get_seq_length"): + cache_length = past_key_values.get_seq_length() + past_length = getattr(past_key_values, 'seen_tokens', + cache_length) + else: + cache_length = past_length = past_key_values[0][0].shape[2] + + # Keep only the unprocessed tokens: + # 1 - If the length of the attention_mask exceeds the length of input_ids, then we are in a setting where + # some of the inputs are exclusively passed as part of the cache (e.g. when passing input_embeds as + # input) + if attention_mask is not None and attention_mask.shape[ + 1] > input_ids.shape[1]: + input_ids = input_ids[:, -(attention_mask.shape[1] - + past_length):] + # 2 - If the past_length is smaller than input_ids', then input_ids holds all input tokens. We can discard + # input_ids based on the past_length. + elif past_length < input_ids.shape[1]: + input_ids = input_ids[:, past_length:] + # 3 - Otherwise (past_length >= input_ids.shape[1]), let's assume input_ids only has unprocessed tokens. + elif self.config.media_placeholder_token_id in input_ids: + input_ids = input_ids[:, input_ids.shape[1] - 1:] + # If the cache has seen more tokens than it can hold, then the cache has a size limit. Let's discard the + # older attention values, as their corresponding values are not part of the input. + if cache_length < past_length and attention_mask is not None: + attention_mask = attention_mask[:, -(cache_length + + input_ids.shape[1]):] + + position_ids = kwargs.get("position_ids", None) + if attention_mask is not None and position_ids is None: + # create position_ids on the fly for batch generation + position_ids = attention_mask.long().cumsum(-1) - 1 + position_ids.masked_fill_(attention_mask == 0, 1) + if past_key_values: + position_ids = position_ids[:, -input_ids.shape[1]:] + + # if `inputs_embeds` are passed, we only want to use them in the 1st generation step + if inputs_embeds is not None and past_key_values is None: + model_inputs = {"inputs_embeds": inputs_embeds} + else: + model_inputs = {"input_ids": input_ids} + + model_inputs.update({ + "position_ids": position_ids, + "past_key_values": past_key_values, + "use_cache": kwargs.get("use_cache"), + "attention_mask": attention_mask, + "pixel_values": pixel_values, + "grid_thws": grid_thws, + }) + return model_inputs + + def _reorder_cache(self, *args, **kwargs): + return self.language_model._reorder_cache(*args, **kwargs) diff --git a/modeling_kimi_linear.py b/modeling_kimi_linear.py new file mode 100644 index 0000000000000000000000000000000000000000..b8c41e8bfce768d74d8da3a37e693f5ee43876a0 --- /dev/null +++ b/modeling_kimi_linear.py @@ -0,0 +1,1314 @@ +# coding=utf-8 +# Copyright 2025-2026 The Moonshot AI Team, DeepSeek-AI, and HuggingFace Inc. team. All rights reserved. +# +# The multi-head latent attention, MoE gating and sparse MoE block in this file are +# adapted from DeepSeek-V3 (DeepSeek-V3/modeling_deepseek.py). They have been +# extensively modified and extended for the Kimi-Linear architecture. +# +# Licensing Information: +# - Code adapted from DeepSeek-V3 (DeepSeek-V3/modeling_deepseek.py) is licensed under the Apache License, Version 2.0. +# - Other parts of the code are licensed under the Kimi K3 License (see the LICENSE file in this repository). +# +# Apache License, Version 2.0: +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +import math +from collections.abc import Callable +from typing import Any + +import torch +import torch.nn.functional as F +import transformers +from einops import rearrange +from packaging import version +from torch import nn +from transformers.activations import ACT2FN +from transformers.cache_utils import Cache +from transformers.generation import GenerationMixin +from transformers.masking_utils import create_causal_mask +from transformers.modeling_flash_attention_utils import FlashAttentionKwargs +from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast +from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel +from transformers.processing_utils import Unpack +from transformers.pytorch_utils import ALL_LAYERNORM_LAYERS +from transformers.utils import TransformersKwargs, auto_docstring, can_return_tuple, logging +from transformers.utils.generic import OutputRecorder, check_model_inputs + +try: + from fla.modules import FusedRMSNormGated, ShortConvolution + from fla.ops.kda import chunk_kda, fused_recurrent_kda + # from fla.ops.kda.gate import fused_kda_gate # deprecated, gate is now computed inside chunk_kda/fused_recurrent_kda + from fla.ops.utils.index import prepare_cu_seqlens_from_mask, prepare_lens_from_mask + from fla.utils import tensor_cache +except ImportError: + raise ImportError("Plese run `pip install -U fla-core`") + +from .configuration_kimi_k3 import KimiLinearConfig + +assert version.parse(transformers.__version__) >= version.parse("4.56.0"), \ + "Please upgrade transformers to >= 4.56.0" + +logger = logging.get_logger(__name__) + + +# Register Moonshot-specific activation functions +class SituAndMul(nn.Module): + """ + SituAndMul activation: beta * tanh(gate / beta) * sigmoid(gate) * up + When linear_beta is set, up is also transformed by linear_beta * tanh(up / linear_beta). + """ + + def __init__(self, beta: float = 1.0, linear_beta: float | None = None): + super().__init__() + self.beta = beta + self.linear_beta = linear_beta + + def forward(self, x: torch.Tensor) -> torch.Tensor: + d = x.shape[-1] // 2 + gate = x[..., :d].to(torch.float32) + up = x[..., d:].to(torch.float32) + situ_a = self.beta * torch.tanh(gate / self.beta) * torch.sigmoid(gate) + if self.linear_beta is not None: + up = self.linear_beta * torch.tanh(up / self.linear_beta) + return (situ_a * up).to(x.dtype) + + +ACT2FN["situ"] = SituAndMul + + +def _get_situ_activation_params(config: KimiLinearConfig): + beta = getattr(config, "activation_situ_beta", None) + linear_beta = getattr(config, "activation_situ_linear_beta", None) + return beta or 1.0, linear_beta + + +def index_first_axis(x, indices): + return x[indices] + + +@tensor_cache +def get_unpad_data( + attention_mask: torch.Tensor, +) -> tuple[torch.Tensor, torch.Tensor, int]: + lens = prepare_lens_from_mask(attention_mask) + indices = torch.nonzero(attention_mask.flatten(), as_tuple=False).flatten() + max_seqlen_in_batch = lens.max().item() + cu_seqlens = prepare_cu_seqlens_from_mask(attention_mask) + return indices, cu_seqlens, max_seqlen_in_batch + + +def pad_input( + hidden_states: torch.Tensor, + indices: torch.LongTensor, + batch_size: int, + seq_len: int, +) -> torch.Tensor: + out = hidden_states.new_zeros((batch_size * seq_len, *hidden_states.shape[1:])) + out[indices] = hidden_states + return out.view(batch_size, seq_len, *hidden_states.shape[1:]) + + +class KimiDynamicCache: + """ + Dynamic cache for Kimi model. + Inspired by Qwen3-Next + """ + is_compileable = False + + def __init__(self, config: KimiLinearConfig): + super().__init__() + self.config = config + + if config.linear_attn_config is not None: + self.layer_types = [] + for i in range(config.num_hidden_layers): + if config.is_kda_layer(i): + self.layer_types.append("linear_attention") + else: + self.layer_types.append("full_attention") + else: + self.layer_types = ["full_attention"] * config.num_hidden_layers + + self.transformer_layers = [ + i for i in range(config.num_hidden_layers) if self.layer_types[i] == "full_attention" + ] + + linear_layers = [i for i in range( + config.num_hidden_layers) if self.layer_types[i] == "linear_attention"] + self.last_linear_layer = linear_layers[-1] if linear_layers else -1 + + self.conv_states = [None for _ in range(config.num_hidden_layers)] + self.recurrent_states = [None for _ in range(config.num_hidden_layers)] + self.key_cache = [None for _ in range(config.num_hidden_layers)] + self.value_cache = [None for _ in range(config.num_hidden_layers)] + + def __len__(self): + return len(self.layer_types) + + def update( + self, + key_states: torch.Tensor, + value_states: torch.Tensor, + layer_idx: int, + cache_kwargs: dict[str, Any] | None = None, + ) -> tuple[torch.Tensor, torch.Tensor]: + if self.key_cache[layer_idx] is None: + self.key_cache[layer_idx] = key_states + self.value_cache[layer_idx] = value_states + else: + self.key_cache[layer_idx] = torch.cat( + [self.key_cache[layer_idx], key_states], dim=2) + self.value_cache[layer_idx] = torch.cat( + [self.value_cache[layer_idx], value_states], dim=2) + + return self.key_cache[layer_idx], self.value_cache[layer_idx] + + def reorder_cache(self, beam_idx: torch.LongTensor): + """Reorders the cache for beam search, given the selected beam indices.""" + for layer_idx in range(len(self.key_cache)): + if self.key_cache[layer_idx] is not None: + device = self.key_cache[layer_idx].device + beam_idx = beam_idx.to(device) + self.key_cache[layer_idx] = self.key_cache[layer_idx].index_select( + 0, beam_idx) + self.value_cache[layer_idx] = self.value_cache[layer_idx].index_select( + 0, beam_idx) + + if self.conv_states[layer_idx] is not None: + device = self.conv_states[layer_idx][0].device + beam_idx = beam_idx.to(device) + q_conv, k_conv, v_conv = self.conv_states[layer_idx] + self.conv_states[layer_idx] = ( + q_conv.index_select(0, beam_idx), + k_conv.index_select(0, beam_idx), + v_conv.index_select(0, beam_idx), + ) + self.recurrent_states[layer_idx] = self.recurrent_states[layer_idx].index_select( + 0, beam_idx) + + def get_seq_length(self, layer_idx: int | None = 0) -> int: + """Returns the sequence length of the cached states. A layer index can be optionally passed.""" + # take any layer that contains cache and not empty tensor + layer_idx = self.transformer_layers[0] if layer_idx not in self.transformer_layers else layer_idx + if len(self.key_cache) <= layer_idx or self.key_cache[layer_idx] is None: + return 0 + return self.key_cache[layer_idx].shape[-2] + + def get_mask_sizes(self, cache_position: torch.Tensor, layer_idx: int) -> tuple[int, int]: + """ + Return a tuple (kv_length, kv_offset) corresponding to the length and offset that will be returned for + the given layer at `layer_idx`. + The masks are then prepared according to the given lengths (kv_length, kv_offset) and patterns for each layer. + """ + kv_offset = 0 + query_length = cache_position.shape[0] + past_seen_tokens = self.get_seq_length(layer_idx) + kv_length = query_length + past_seen_tokens + return kv_length, kv_offset + + @property + def has_previous_state(self): + """We have a previous state if the last linear (conv) layer was already updated.""" + if self.last_linear_layer == -1: + return False + return self.conv_states[self.last_linear_layer] is not None + + +class KimiRMSNorm(nn.Module): + def __init__(self, hidden_size, eps=1e-6): + super().__init__() + self.weight = nn.Parameter(torch.ones(hidden_size)) + self.variance_epsilon = eps + + def forward(self, hidden_states): + dtype = hidden_states.dtype + x = hidden_states.float() + x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.variance_epsilon) + return self.weight * x.to(dtype) + + +ALL_LAYERNORM_LAYERS.append(KimiRMSNorm) + + +class KimiBlockSparseMLP(nn.Module): + def __init__(self, config: KimiLinearConfig, hidden_size=None, intermediate_size=None): + super().__init__() + self.config = config + self.ffn_dim = config.intermediate_size if intermediate_size is None else intermediate_size + self.hidden_dim = config.hidden_size if hidden_size is None else hidden_size + + self.w1 = nn.Linear(self.hidden_dim, self.ffn_dim, bias=False) # gate + self.w2 = nn.Linear(self.ffn_dim, self.hidden_dim, bias=False) # down + self.w3 = nn.Linear(self.hidden_dim, self.ffn_dim, bias=False) # up + + if config.hidden_act == "situ": + beta, linear_beta = _get_situ_activation_params(config) + self.act_fn = SituAndMul( + beta=beta, + linear_beta=linear_beta, + ) + else: + self.act_fn = ACT2FN[config.hidden_act] + + def forward(self, hidden_states): + if self.config.hidden_act == "situ": + gate_up = torch.cat([self.w1(hidden_states), self.w3(hidden_states)], dim=-1) + current_hidden_states = self.act_fn(gate_up) + else: + current_hidden_states = self.act_fn( + self.w1(hidden_states)) * self.w3(hidden_states) + current_hidden_states = self.w2(current_hidden_states) + return current_hidden_states + + +class KimiMLP(nn.Module): + def __init__(self, config: KimiLinearConfig, hidden_size=None, intermediate_size=None): + super().__init__() + self.config = config + self.hidden_size = config.hidden_size if hidden_size is None else hidden_size + self.intermediate_size = config.intermediate_size if intermediate_size is None else intermediate_size + self.gate_proj = nn.Linear( + self.hidden_size, self.intermediate_size, bias=False) + self.up_proj = nn.Linear( + self.hidden_size, self.intermediate_size, bias=False) + self.down_proj = nn.Linear( + self.intermediate_size, self.hidden_size, bias=False) + if config.hidden_act == "situ": + beta, linear_beta = _get_situ_activation_params(config) + self.act_fn = SituAndMul( + beta=beta, + linear_beta=linear_beta, + ) + else: + self.act_fn = ACT2FN[config.hidden_act] + + def forward(self, x): + if self.config.hidden_act == "situ": + gate_up = torch.cat([self.gate_proj(x), self.up_proj(x)], dim=-1) + down_proj = self.down_proj(self.act_fn(gate_up)) + else: + down_proj = self.down_proj(self.act_fn( + self.gate_proj(x)) * self.up_proj(x)) + return down_proj + + +def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: + """Expand the key/value heads from `num_key_value_heads` to `num_attention_heads`.""" + if n_rep == 1: + return hidden_states + return torch.repeat_interleave(hidden_states, dim=1, repeats=n_rep) + + +def eager_attention_forward( + module: nn.Module, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attention_mask: torch.Tensor | None, + scaling: float, + dropout: float = 0.0, + **kwargs: Unpack[TransformersKwargs], +): + key = repeat_kv(key, module.num_key_value_groups) + value = repeat_kv(value, module.num_key_value_groups) + + scores = torch.einsum("bhqd,bhkd->bhqk", query, key) * scaling + if attention_mask is not None: + scores = scores + attention_mask[:, :, :, : key.shape[-2]] + + probs = F.softmax(scores, dim=-1, dtype=torch.float32).to(query.dtype) + probs = F.dropout(probs, p=dropout, training=module.training) + out = torch.einsum("bhqk,bhkd->bhqd", probs, value).transpose(1, 2).contiguous() + + return out, probs + + +class KimiMLAAttention(nn.Module): + """ + Multi-Latent Attention adapted from deepseek-v3 + """ + + def __init__(self, config: KimiLinearConfig, layer_idx: int): + nn.Module.__init__(self) + self.config = config + self.layer_idx = layer_idx + self.hidden_size = config.hidden_size + self.num_heads = config.num_attention_heads + self.num_key_value_heads = config.num_key_value_heads + self.num_key_value_groups = self.num_heads // self.num_key_value_heads + + self.attention_dropout = getattr(config, "attention_dropout", 0.0) + + try: + self.q_lora_rank = config.q_lora_rank + self.qk_rope_head_dim = config.qk_rope_head_dim + self.kv_lora_rank = config.kv_lora_rank + self.v_head_dim = config.v_head_dim + self.qk_nope_head_dim = config.qk_nope_head_dim + self.q_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim + self.use_nope = config.mla_use_nope + self.scaling = self.q_head_dim ** (-0.5) + except Exception as e: + raise ValueError( + f"Kimi MLA config is not found or not properly formatted: {e}") + + if self.q_lora_rank is not None: + self.q_a_proj = nn.Linear( + self.hidden_size, self.q_lora_rank, bias=False, + ) + self.q_a_layernorm = KimiRMSNorm(self.q_lora_rank) + self.q_b_proj = nn.Linear( + self.q_lora_rank, + self.num_heads * self.q_head_dim, + bias=False, + ) + else: + self.q_proj = nn.Linear( + self.hidden_size, self.num_heads * self.q_head_dim, bias=False, + ) + self.kv_a_proj_with_mqa = nn.Linear( + self.hidden_size, + self.kv_lora_rank + self.qk_rope_head_dim, + bias=False, + ) + self.kv_a_layernorm = KimiRMSNorm(self.kv_lora_rank) + self.kv_b_proj = nn.Linear( + self.kv_lora_rank, + self.num_heads + * (self.q_head_dim - self.qk_rope_head_dim + self.v_head_dim), + bias=False, + ) + self.o_proj = nn.Linear( + self.num_heads * self.v_head_dim, + self.hidden_size, + bias=False, + ) + self.is_causal = True + assert self.use_nope + + self.use_output_gate = getattr(config, "mla_use_output_gate", False) + if self.use_output_gate: + projection_size = self.num_heads * self.v_head_dim + self.g_proj = nn.Linear(self.hidden_size, projection_size, bias=False) + + self.rotary_emb = None + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Cache | None = None, + **kwargs, + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + batch_size, seq_length = hidden_states.shape[:-1] + query_shape = (batch_size, seq_length, -1, self.q_head_dim) + key_shape = (batch_size, seq_length, -1, + self.qk_nope_head_dim + self.v_head_dim) + + if self.q_lora_rank is not None: + q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states))) + else: + q_states = self.q_proj(hidden_states) + q_states = q_states.view(query_shape).transpose(1, 2) + q_pass, q_rot = torch.split( + q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + + compressed_kv = self.kv_a_proj_with_mqa(hidden_states) + k_pass, k_rot = torch.split( + compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + + k_pass = self.kv_b_proj(self.kv_a_layernorm( + k_pass)).view(key_shape).transpose(1, 2) + k_pass, value_states = torch.split( + k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) + + k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim) + + k_rot = k_rot.expand(*k_pass.shape[:-1], -1) + + query_states = torch.cat((q_pass, q_rot), dim=-1) + key_states = torch.cat((k_pass, k_rot), dim=-1) + + if past_key_values is not None: + key_states, value_states = past_key_values.update( + key_states, value_states, self.layer_idx) + + if self.config._attn_implementation == "flash_attention_2" and self.q_head_dim != self.v_head_dim: + value_states = F.pad( + value_states, [0, self.q_head_dim - self.v_head_dim]) + + attention_interface: Callable = eager_attention_forward + if self.config._attn_implementation != "eager": + attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] + + attn_output, _ = attention_interface( + self, + query_states, + key_states, + value_states, + attention_mask, + dropout=0.0 if not self.training else self.attention_dropout, + scaling=self.scaling, + **kwargs, + ) + + if self.config._attn_implementation == "flash_attention_2" and self.q_head_dim != self.v_head_dim: + attn_output = attn_output[:, :, :, : self.v_head_dim] + + attn_output = attn_output.reshape( + batch_size, seq_length, -1).contiguous() + if self.use_output_gate: + g = self.g_proj(hidden_states).sigmoid() + attn_output = attn_output * g + attn_output = self.o_proj(attn_output) + return attn_output + + +class KimiDeltaAttention(nn.Module): + def __init__(self, config: KimiLinearConfig, layer_idx: int): + super().__init__() + self.config = config + self.mode = "chunk" + + self.hidden_size = config.hidden_size + self.conv_size = config.linear_attn_config["short_conv_kernel_size"] + self.head_dim = config.linear_attn_config["head_dim"] + self.num_heads = config.linear_attn_config["num_heads"] + self.head_k_dim = self.head_dim + self.num_k_heads = self.num_heads + + self.layer_idx = layer_idx + + assert self.mode in [ + 'chunk', 'fused_recurrent'], f"Not supported mode `{self.mode}`." + + projection_k_size = self.head_k_dim * self.num_k_heads + projection_size = self.head_dim * self.num_heads + + self.q_proj = nn.Linear( + self.hidden_size, projection_k_size, bias=False) + self.k_proj = nn.Linear( + self.hidden_size, projection_k_size, bias=False) + self.v_proj = nn.Linear(self.hidden_size, projection_size, bias=False) + + self.q_conv1d = ShortConvolution( + hidden_size=projection_k_size, + kernel_size=self.conv_size, + activation='silu', + ) + self.k_conv1d = ShortConvolution( + hidden_size=projection_k_size, + kernel_size=self.conv_size, + activation='silu', + ) + self.v_conv1d = ShortConvolution( + hidden_size=projection_size, + kernel_size=self.conv_size, + activation='silu', + ) + + self.A_log = torch.nn.Parameter(torch.log(torch.empty( + self.num_heads, dtype=torch.float32).uniform_(1, 16))) + + self.f_a_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False) + self.f_b_proj = nn.Linear(self.head_dim, projection_size, bias=False) + + self.dt_bias = nn.Parameter( + torch.empty(projection_size, dtype=torch.float32)) + + self.b_proj = nn.Linear(self.hidden_size, self.num_heads, bias=False) + + self.use_full_rank_gate = config.linear_attn_config.get("use_full_rank_gate", False) + self.gate_lower_bound = config.linear_attn_config.get("gate_lower_bound", None) + if self.use_full_rank_gate: + self.g_proj = nn.Linear(self.hidden_size, projection_size, bias=False) + else: + self.g_a_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False) + self.g_b_proj = nn.Linear(self.head_dim, projection_size, bias=False) + + self.o_norm = FusedRMSNormGated( + self.head_dim, eps=config.rms_norm_eps, activation='sigmoid') + self.o_proj = nn.Linear(projection_size, self.hidden_size, bias=False) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + cache_params: KimiDynamicCache | None = None, + **kwargs: Unpack[dict], + ) -> tuple[torch.Tensor, torch.Tensor | None, Cache | None]: + if attention_mask is not None: + if attention_mask.dim() != 2: + attention_mask = kwargs.get("padding_mask") + + if attention_mask is not None and attention_mask.dim() != 2: + raise ValueError( + "attention_mask must be a 0-1 matrix of shape [batch_size, seq_len] " + "(0 = padding). 3D masks are not supported here.", + ) + use_cache = cache_params is not None + batch_size, q_len, _ = hidden_states.shape + mode = 'fused_recurrent' if use_cache and q_len == 1 else self.mode + if self.training: + assert mode == 'chunk', "Only chunk mode is supported in training." + + cu_seqlens = kwargs.get('cu_seqlens') + indices = None + if attention_mask is not None: + indices, cu_seqlens, _ = get_unpad_data(attention_mask[:, -q_len:]) + hidden_states = index_first_axis( + rearrange(hidden_states, "b s ... -> (b s) ..."), indices).unsqueeze(0) + + conv_state_q, conv_state_k, conv_state_v = None, None, None + recurrent_state = None + if cache_params is not None: + if cache_params.conv_states[self.layer_idx] is not None: + conv_state_q, conv_state_k, conv_state_v = cache_params.conv_states[ + self.layer_idx] + recurrent_state = cache_params.recurrent_states[self.layer_idx] + + q_proj_states = self.q_proj(hidden_states) + k_proj_states = self.k_proj(hidden_states) + v_proj_states = self.v_proj(hidden_states) + q, conv_state_q = self.q_conv1d( + x=q_proj_states, + cache=conv_state_q, + output_final_state=use_cache, + cu_seqlens=cu_seqlens, + ) + k, conv_state_k = self.k_conv1d( + x=k_proj_states, + cache=conv_state_k, + output_final_state=use_cache, + cu_seqlens=cu_seqlens, + ) + v, conv_state_v = self.v_conv1d( + x=v_proj_states, + cache=conv_state_v, + output_final_state=use_cache, + cu_seqlens=cu_seqlens, + ) + g = self.f_b_proj(self.f_a_proj(hidden_states)) + g = rearrange(g, '... (h d) -> ... h d', d=self.head_dim) + beta = self.b_proj(hidden_states).float() + + q, k = map(lambda x: rearrange( + x, '... (h d) -> ... h d', d=self.head_k_dim), (q, k)) + v = rearrange(v, '... (h d) -> ... h d', d=self.head_dim) + + if mode == 'chunk': + o, recurrent_state = chunk_kda( + q=q, + k=k, + v=v, + g=g, + beta=beta, + A_log=self.A_log, + dt_bias=self.dt_bias, + initial_state=recurrent_state, + output_final_state=True, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + use_beta_sigmoid_in_kernel=True, + safe_gate=self.gate_lower_bound is not None, + lower_bound=self.gate_lower_bound, + transpose_state_layout=True, + cu_seqlens=cu_seqlens, + ) + else: + o, recurrent_state = fused_recurrent_kda( + q=q, + k=k, + v=v, + g=g, + beta=beta, + A_log=self.A_log, + dt_bias=self.dt_bias, + initial_state=recurrent_state, + output_final_state=True, + use_qk_l2norm_in_kernel=True, + use_gate_in_kernel=True, + use_beta_sigmoid_in_kernel=True, + lower_bound=self.gate_lower_bound, + transpose_state_layout=True, + cu_seqlens=cu_seqlens, + ) + if cache_params is not None: + cache_params.recurrent_states[self.layer_idx] = recurrent_state + cache_params.conv_states[self.layer_idx] = ( + conv_state_q, conv_state_k, conv_state_v) + + if self.use_full_rank_gate: + g = self.g_proj(hidden_states) + else: + g = self.g_b_proj(self.g_a_proj(hidden_states)) + g = rearrange(g, '... (h d) -> ... h d', d=self.head_dim) + o = self.o_norm(o, g) + + o = rearrange(o, 'b t h d -> b t (h d)') + o = self.o_proj(o) + if attention_mask is not None: + o = pad_input(o.squeeze(0), indices, batch_size, q_len) + + return o + + +class KimiMoEGate(nn.Module): + """ + MoEGate adapted from Deepseek-V3. + Parameter correspondences: + num_experts -> n_routed_experts + num_experts_per_token -> num_experts_per_tok + num_expert_group -> n_group + moe_router_activation_func -> scoring_func + """ + + def __init__(self, config: KimiLinearConfig): + super().__init__() + self.config = config + self.top_k = config.num_experts_per_token + self.num_experts = config.num_experts + self.routed_scaling_factor = config.routed_scaling_factor + self.moe_router_activation_func = config.moe_router_activation_func + self.num_expert_group = getattr(config, "num_expert_group", 1) + self.topk_group = getattr(config, "topk_group", 1) + + # topk selection algorithm + self.moe_renormalize = config.moe_renormalize + self.gating_dim = config.hidden_size + self.weight = nn.Parameter( + torch.empty((self.num_experts, self.gating_dim)), + ) + + self.e_score_correction_bias = nn.Parameter( + torch.empty(self.num_experts), + ) + self.reset_parameters() + + def reset_parameters(self) -> None: + import torch.nn.init as init + + init.kaiming_uniform_(self.weight, a=math.sqrt(5)) + + def forward(self, hidden_states): + bsz, seq_len, h = hidden_states.shape + # compute gating score + hidden_states = hidden_states.view(-1, h) + logits = F.linear( + hidden_states.type(torch.float32), self.weight.type( + torch.float32), None, + ) + if self.moe_router_activation_func == "sigmoid": + scores = logits.sigmoid() + elif self.moe_router_activation_func == "softmax": + scores = logits.softmax(dim=1) + else: + raise NotImplementedError( + f"insupportable scoring function for MoE gating: {self.moe_router_activation_func}", + ) + + # select top-k experts + assert not self.training + scores = scores.view(bsz * seq_len, -1) + scores_for_choice = scores + self.e_score_correction_bias.unsqueeze(0) + if self.num_expert_group > 1 and self.num_expert_group > self.topk_group: + group_scores = ( + scores_for_choice.view( + bsz * seq_len, self.num_expert_group, -1).topk(2, dim=-1)[0].sum(dim=-1) + ) # [n, num_expert_group] + group_idx = torch.topk( + group_scores, k=self.topk_group, dim=-1, sorted=False, + )[ + 1 + ] # [n, top_k_group] + group_mask = torch.zeros_like(group_scores) # [n, num_expert_group] + group_mask.scatter_(1, group_idx, 1) # [n, num_expert_group] + score_mask = ( + group_mask.unsqueeze(-1) + .expand( + bsz * seq_len, self.num_expert_group, self.num_experts // self.num_expert_group, + ) + .reshape(bsz * seq_len, -1) + ) # [n, e] + tmp_scores = scores_for_choice.masked_fill( + ~score_mask.bool(), float("-inf")) # [n, e] + else: + tmp_scores = scores_for_choice + _, topk_idx = torch.topk( + tmp_scores, k=self.top_k, dim=-1, sorted=False, + ) + topk_weight = scores.gather(1, topk_idx) + + # norm gate to sum 1 + if self.top_k > 1 and self.moe_renormalize: + denominator = topk_weight.sum(dim=-1, keepdim=True) + 1e-20 + topk_weight = topk_weight / denominator + # must multiply the scaling factor + topk_weight = topk_weight * self.routed_scaling_factor + + return topk_idx, topk_weight + + +class KimiSparseMoeBlock(nn.Module): + """ + Adapted from Deepseek-V3's MOE implementation + The namings are consistent with Kimi's version. + """ + + def __init__(self, config: KimiLinearConfig): + super().__init__() + self.config = config + self.hidden_dim = config.hidden_size + self.num_experts = config.num_experts + self.top_k = config.num_experts_per_token + self.moe_renormalize = config.moe_renormalize + + self.use_latent_moe = getattr(config, "routed_expert_hidden_size", None) is not None + self.moe_hidden_size = ( + config.routed_expert_hidden_size + if self.use_latent_moe else config.hidden_size + ) + self.latent_moe_use_norm = getattr(config, "latent_moe_use_norm", False) + + self.ep_size = 1 + self.experts_per_rank = config.num_experts + self.ep_rank = 0 + self.experts = nn.ModuleList( + [ + KimiBlockSparseMLP( + config, + hidden_size=self.moe_hidden_size, + intermediate_size=config.moe_intermediate_size, + ) + for _ in range(config.num_experts) + ], + ) + self.gate = KimiMoEGate(config) + if config.num_shared_experts is not None: + intermediate_size = config.moe_intermediate_size * config.num_shared_experts + self.shared_experts = KimiMLP( + config=config, intermediate_size=intermediate_size, + ) + + if self.use_latent_moe: + self.routed_expert_down_proj = nn.Linear( + config.hidden_size, self.moe_hidden_size, bias=False, + ) + self.routed_expert_up_proj = nn.Linear( + self.moe_hidden_size, config.hidden_size, bias=False, + ) + if self.latent_moe_use_norm: + self.routed_expert_norm = KimiRMSNorm( + self.moe_hidden_size, eps=config.rms_norm_eps, + ) + + def forward(self, hidden_states): + identity = hidden_states + orig_shape = hidden_states.shape + topk_idx, topk_weight = self.gate(hidden_states) + hidden_states = hidden_states.view(-1, hidden_states.shape[-1]) + + if self.use_latent_moe: + hidden_states = self.routed_expert_down_proj(hidden_states) + + if not self.training: + y = self.moe_infer(hidden_states, topk_idx, topk_weight) + else: + raise NotImplementedError("Training mode is not supported in KimiSparseMoeBlock") + + if self.use_latent_moe: + if self.latent_moe_use_norm: + y = self.routed_expert_norm(y) + y = self.routed_expert_up_proj(y) + + y = y.view(*orig_shape) + + if self.config.num_shared_experts is not None: + y = y + self.shared_experts(identity) + return y + + @torch.no_grad() + def moe_infer(self, x, topk_ids, topk_weight): + cnts = topk_ids.new_zeros((topk_ids.shape[0], len(self.experts))) + cnts.scatter_(1, topk_ids, 1) + tokens_per_expert = cnts.sum(dim=0) + idxs = topk_ids.view(-1).argsort() + sorted_tokens = x[idxs // topk_ids.shape[1]] + + tokens_per_expert = tokens_per_expert.cpu().numpy() + + outputs = [] + start_idx = 0 + for i, num_tokens in enumerate(tokens_per_expert): + end_idx = start_idx + num_tokens + if num_tokens == 0: + continue + expert = self.experts[i + self.ep_rank * self.experts_per_rank] + tokens_for_this_expert = sorted_tokens[start_idx:end_idx] + expert_out = expert(tokens_for_this_expert) + outputs.append(expert_out) + start_idx = end_idx + + outs = torch.cat(outputs, dim=0) if len( + outputs) else sorted_tokens.new_empty(0) + + new_x = torch.empty_like(outs) + new_x[idxs] = outs + final_out = ( + new_x.view(*topk_ids.shape, -1) + .type(topk_weight.dtype) + .mul_(topk_weight.unsqueeze(dim=-1)) + .sum(dim=1) + .type(new_x.dtype) + ) + return final_out + + +class KimiDecoderLayer(nn.Module): + def __init__(self, config: KimiLinearConfig, layer_idx: int): + super().__init__() + self.hidden_size = config.hidden_size + self.config = config + self.layer_idx = layer_idx + if config.is_kda_layer(layer_idx): + self.is_linear_attn = True + self.self_attn = KimiDeltaAttention( + config=config, layer_idx=layer_idx) + elif config.is_mla: + self.is_linear_attn = False + self.self_attn = KimiMLAAttention( + config=config, layer_idx=layer_idx) + else: + raise NotImplementedError + if ( + config.num_experts is not None + and layer_idx >= config.first_k_dense_replace + and layer_idx % getattr(config, "moe_layer_freq", 1) == 0 + ): + self.block_sparse_moe = KimiSparseMoeBlock(config) + else: + self.mlp = KimiMLP(config) + self.input_layernorm = KimiRMSNorm( + config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = KimiRMSNorm( + config.hidden_size, eps=config.rms_norm_eps) + + # Attention residual + self.use_attn_residuals = getattr(config, "attn_res_block_size", None) is not None + if self.use_attn_residuals: + self.attn_res_block_size = config.attn_res_block_size + self.self_attention_res_norm = KimiRMSNorm( + config.hidden_size, eps=config.rms_norm_eps) + self.mlp_res_norm = KimiRMSNorm( + config.hidden_size, eps=config.rms_norm_eps) + self.self_attention_res_proj = nn.Linear( + config.hidden_size, 1, bias=False) + self.mlp_res_proj = nn.Linear( + config.hidden_size, 1, bias=False) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: tuple[torch.Tensor] | None = None, + output_attentions: bool | None = False, + use_cache: bool | None = False, + block_residual: torch.Tensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ): + if self.use_attn_residuals: + return self._forward_attn_residual( + hidden_states, attention_mask, position_ids, + past_key_values, output_attentions, use_cache, + block_residual, **kwargs) + + residual = hidden_states + + hidden_states = self.input_layernorm(hidden_states) + + # Self Attention + if self.is_linear_attn is False: + hidden_states = self.self_attn( + hidden_states=hidden_states, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + output_attentions=output_attentions, + use_cache=use_cache, + **kwargs, + ) + else: + hidden_states = self.self_attn( + hidden_states=hidden_states, + attention_mask=attention_mask, + cache_params=past_key_values, + output_attentions=output_attentions, + use_cache=use_cache, + **kwargs, + ) + hidden_states = residual + hidden_states + + # Fully Connected + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + if hasattr(self, "block_sparse_moe"): + hidden_states = self.block_sparse_moe(hidden_states) + else: + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + + return hidden_states + + def _forward_attn_residual( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: tuple[torch.Tensor] | None = None, + output_attentions: bool | None = False, + use_cache: bool | None = False, + block_residual: torch.Tensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ): + batch_size, seq_len, hidden_size = hidden_states.shape + prefix_sum = hidden_states + + if block_residual is not None and block_residual.shape[1] > 0: + hidden_states = _apply_attn_res( + prefix_sum.view(-1, hidden_size), + block_residual, + self.self_attention_res_proj, + self.self_attention_res_norm, + ).view(batch_size, seq_len, hidden_size) + + if self.layer_idx % self.attn_res_block_size == 0: + block_residual = torch.cat( + [block_residual, prefix_sum.view(-1, hidden_size).unsqueeze(1)], dim=1) + prefix_sum = None + + hidden_states = self.input_layernorm(hidden_states) + + # Self Attention + if self.is_linear_attn is False: + hidden_states = self.self_attn( + hidden_states=hidden_states, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + output_attentions=output_attentions, + use_cache=use_cache, + **kwargs, + ) + else: + hidden_states = self.self_attn( + hidden_states=hidden_states, + attention_mask=attention_mask, + cache_params=past_key_values, + output_attentions=output_attentions, + use_cache=use_cache, + **kwargs, + ) + + if prefix_sum is not None: + prefix_sum = prefix_sum + hidden_states + else: + prefix_sum = hidden_states + + hidden_states = _apply_attn_res( + prefix_sum.view(-1, hidden_size), + block_residual, + self.mlp_res_proj, + self.mlp_res_norm, + ).view(batch_size, seq_len, hidden_size) + + hidden_states = self.post_attention_layernorm(hidden_states) + if hasattr(self, "block_sparse_moe"): + hidden_states = self.block_sparse_moe(hidden_states) + else: + hidden_states = self.mlp(hidden_states) + + if prefix_sum is None: + prefix_sum = hidden_states + else: + prefix_sum = prefix_sum + hidden_states + + return prefix_sum, block_residual + + +class KimiPreTrainedModel(PreTrainedModel): + config_class = KimiLinearConfig + base_model_prefix = "model" + supports_gradient_checkpointing = True + _no_split_modules = ["KimiDecoderLayer"] + _skip_keys_device_placement = "past_key_values" + _supports_flash_attn_2 = True + _can_record_outputs = { + "router_logits": OutputRecorder(KimiBlockSparseMLP, index=1), + "hidden_states": KimiDecoderLayer, + "attentions": KimiMLAAttention, + } + _is_stateful = True + + def _init_weights(self, module): + std = self.config.initializer_range + if isinstance(module, nn.Linear): + module.weight.data.normal_(mean=0.0, std=std) + if module.bias is not None: + module.bias.data.zero_() + elif isinstance(module, nn.Embedding): + module.weight.data.normal_(mean=0.0, std=std) + if module.padding_idx is not None: + module.weight.data[module.padding_idx].zero_() + + +def _apply_attn_res(prefix_sum, block_residual, proj, norm): + """ + prefix_sum: (num_tokens, hidden_size) + block_residual: (num_tokens, num_blocks, hidden_size) + """ + v = torch.cat((block_residual, prefix_sum.unsqueeze(1)), dim=1) + v_float = v.float() + variance = v_float.pow(2).mean(-1, keepdim=True) + k = v_float * torch.rsqrt(variance + norm.variance_epsilon) + score_weight = norm.weight.float() * proj.weight.squeeze(0).float() + scores = (k * score_weight).sum(-1) + probs = scores.softmax(-1).unsqueeze(1) + hidden_states = torch.matmul(probs, v_float).squeeze(1) + return hidden_states.to(v.dtype) + +class KimiLinearModel(KimiPreTrainedModel): + def __init__(self, config: KimiLinearConfig): + super().__init__(config) + self.padding_idx = config.pad_token_id + self.vocab_size = config.vocab_size + + self.embed_tokens = nn.Embedding( + config.vocab_size, config.hidden_size, self.padding_idx) + self.layers = nn.ModuleList([KimiDecoderLayer( + config, layer_idx) for layer_idx in range(config.num_hidden_layers)]) + self.norm = KimiRMSNorm( + config.hidden_size, eps=config.rms_norm_eps) + + self.use_attn_residuals = getattr(config, "attn_res_block_size", None) is not None + if self.use_attn_residuals: + self.output_attn_res_norm = KimiRMSNorm( + config.hidden_size, eps=config.rms_norm_eps) + self.output_attn_res_proj = nn.Linear( + config.hidden_size, 1, bias=False) + + if getattr(config, "_attn_implementation", None) is not None: + if config._attn_implementation != "flash_attention_2": + logger.warning_once( + f"Ignoring the provided attention implementation {config._attn_implementation}") + logger.warning_once("Using flash_attention_2 backend instead.") + config._attn_implementation = "flash_attention_2" + else: + config._attn_implementation = "flash_attention_2" + + self._use_flash_attention_2 = config._attn_implementation == "flash_attention_2" + self.gradient_checkpointing = False + # Initialize weights and apply final processing + self.post_init() + + def _update_linear_attn_mask(self, attention_mask, cache_position): + """ + NOTE: Left-padding is used for linear attention mask. + No need for zeroing states when + 1. Cached forward + 2. Attending to all inputs + """ + linear_attn_mask = attention_mask + if cache_position[0] > 0 or (attention_mask is not None and torch.all(attention_mask == 1)): + linear_attn_mask = None + return linear_attn_mask + + @check_model_inputs + # @auto_docstring + def forward( + self, + input_ids: torch.LongTensor = None, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Cache | None = None, + inputs_embeds: torch.FloatTensor | None = None, + cache_position: torch.LongTensor | None = None, + use_cache: bool | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> tuple | BaseModelOutputWithPast: + + use_cache = use_cache if use_cache is not None else self.config.use_cache + + if (input_ids is None) and (inputs_embeds is None): + raise ValueError( + "You must specify exactly one of input_ids or inputs_embeds") + + # Get inputs_embeds + if inputs_embeds is None: + inputs_embeds = self.embed_tokens(input_ids) + + if use_cache and past_key_values is None: + past_key_values = KimiDynamicCache(config=self.config) + + if cache_position is None: + past_seen_tokens = past_key_values.get_seq_length( + ) if past_key_values is not None else 0 + cache_position: torch.Tensor = torch.arange( + past_seen_tokens, past_seen_tokens + inputs_embeds.shape[1], device=inputs_embeds.device, + ) + + if position_ids is None: + position_ids = cache_position.unsqueeze(0) + + causal_mask = create_causal_mask( + config=self.config, + input_embeds=inputs_embeds, + attention_mask=attention_mask, + cache_position=cache_position, + past_key_values=past_key_values, + position_ids=position_ids, + ) + linear_attn_mask = self._update_linear_attn_mask( + attention_mask, cache_position) + + hidden_states = inputs_embeds + if past_key_values is not None: + assert isinstance(past_key_values, KimiDynamicCache) + + block_residual = None + if self.use_attn_residuals: + block_residual = hidden_states.new_zeros( + hidden_states.shape[0] * hidden_states.shape[1], 0, + hidden_states.shape[2]) + + for decoder_layer in self.layers: + layer_mask = linear_attn_mask if decoder_layer.is_linear_attn else causal_mask + + if self.use_attn_residuals: + hidden_states, block_residual = decoder_layer( + hidden_states, + attention_mask=layer_mask, + past_key_values=past_key_values, + cache_position=cache_position, + block_residual=block_residual, + **kwargs, + ) + else: + hidden_states = decoder_layer( + hidden_states, + attention_mask=layer_mask, + past_key_values=past_key_values, + cache_position=cache_position, + **kwargs, + ) + + if self.use_attn_residuals: + hidden_states = self._apply_output_attn_res( + hidden_states, block_residual) + + hidden_states = self.norm(hidden_states) + + return BaseModelOutputWithPast( + last_hidden_state=hidden_states, + past_key_values=past_key_values, + ) + + def _apply_output_attn_res(self, hidden_states, block_residual): + batch_size, seq_len, hidden_size = hidden_states.shape + return _apply_attn_res( + hidden_states.view(-1, hidden_size), + block_residual, + self.output_attn_res_proj, + self.output_attn_res_norm, + ).view(batch_size, seq_len, hidden_size) + + +class KimiLinearForCausalLM(KimiPreTrainedModel, GenerationMixin): + @classmethod + def _supports_default_dynamic_cache(cls) -> bool: + return False + + _tied_weights_keys = ["lm_head.weight"] + + def __init__(self, config): + super().__init__(config) + self.model = KimiLinearModel(config) + self.vocab_size = config.vocab_size + self.lm_head = nn.Linear( + config.hidden_size, config.vocab_size, bias=False) + + # Initialize weights and apply final processing + self.post_init() + + @can_return_tuple + # @auto_docstring + def forward( + self, + input_ids: torch.LongTensor = None, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: list[torch.FloatTensor] | None = None, + inputs_embeds: torch.FloatTensor | None = None, + labels: torch.LongTensor | None = None, + use_cache: bool | None = None, + output_attentions: bool | None = None, + output_hidden_states: bool | None = None, + generation_mode: bool | None = None, + return_dict: bool | None = None, + cache_position: torch.LongTensor | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> tuple | CausalLMOutputWithPast: + r""" + Args: + labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*): + Labels for computing the masked language modeling loss. Indices should either be in `[0, ..., + config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are ignored + (masked), the loss is only computed for the tokens with labels in `[0, ..., config.vocab_size]`. + """ + + output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions + output_hidden_states = ( + output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states + ) + return_dict = return_dict if return_dict is not None else self.config.use_return_dict + + outputs = self.model( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + inputs_embeds=inputs_embeds, + use_cache=use_cache, + output_attentions=output_attentions, + output_hidden_states=output_hidden_states, + return_dict=return_dict, + cache_position=cache_position, + ) + + logits = outputs[0] + if generation_mode: + logits = logits[:, -1:] + logits = self.lm_head(logits) + + loss = None + if labels is not None: + loss = self.loss_function( + logits, labels, self.vocab_size, **kwargs) + + return CausalLMOutputWithPast( + loss=loss, + logits=logits, + past_key_values=outputs.past_key_values, + hidden_states=outputs.hidden_states, + attentions=outputs.attentions, + ) diff --git a/preprocessor_config.json b/preprocessor_config.json new file mode 100644 index 0000000000000000000000000000000000000000..8be28c13c93fb8982f81928d3f4fe80568102286 --- /dev/null +++ b/preprocessor_config.json @@ -0,0 +1,38 @@ +{ + "media_proc_cfg": { + "in_patch_limit": 65536, + "patch_size": 14, + "image_mean": [ + 0.5, + 0.5, + 0.5 + ], + "image_std": [ + 0.5, + 0.5, + 0.5 + ], + "merge_kernel_size": 2, + "fixed_output_tokens": null, + "patch_limit_on_one_side": 512, + "in_patch_limit_each_frame": 16384, + "in_patch_limit_video": 655360, + "sample_fps": 8.0, + "max_num_frames_each_video": null, + "temporal_merge_kernel_size": 4, + "timestamp_mode": "hh:mm:ss.fff", + "transparent_bg_config": { + "pattern": "chessboard", + "chessboard_square_size": 8, + "chessboard_square_on_top_left": true, + "chessboard_white_value": 255, + "chessboard_gray_value": 180 + }, + "transparent_bg_fill_stage": "after_resize", + "config_type": "media_proc.processors.moonvit.MoonViTMediaProcessorConfig" + }, + "auto_map": { + "AutoProcessor": "kimi_k3_processor.KimiK3Processor", + "AutoImageProcessor": "kimi_k3_vision_processing.KimiK3VisionProcessor" + } +} \ No newline at end of file diff --git a/quantize_k3.py b/quantize_k3.py new file mode 100644 index 0000000000000000000000000000000000000000..6e07d75e0d7bcc8f6bb69726acc1501fb27e2812 --- /dev/null +++ b/quantize_k3.py @@ -0,0 +1,1939 @@ +# SPDX-License-Identifier: Apache-2.0 +# SPDX-FileCopyrightText: Copyright contributors to the vLLM project + +"""Quantize Kimi-K3 into data-free MSE-fitted Cubic 2.5-bit mixed precision. + +[Minimum Requirements] + - One or more CUDA GPUs; 32 GiB per worker is the recommended minimum. + +[Python Dependencies] + - Install a CUDA-enabled PyTorch build compatible with the local driver. + - venv/bin/python3 -m pip install torch safetensors + +[Processing Time] + - Approximately 30 minutes on 8xH200; other devices scale with throughput. + +[Example Command] +python -u examples/quantization/quantize_k3.py \ + --source /path/to/Kimi-K3 \ + --output /path/to/Kimi-K3-Cubic-2.5Bit \ + --devices cuda:0,cuda:1,cuda:2,cuda:3,cuda:4,cuda:5,cuda:6,cuda:7 + +Group-local Cubic parameters are selected by a least-squares objective. The +final report presents scale-free NRMSE so loss values are comparable across +bit widths; the square root is reporting-only and does not change candidate +selection. + +Dynamic-A8 carrier correction (round(127*q)/127) is enabled by default. Pass +--disable-a8-correction to fit only the continuous Cubic reconstruction. + +""" + +import argparse +import hashlib +import json +import math +import multiprocessing as mp +import os +import queue +import shutil +import statistics +import time +import traceback +from collections import Counter, defaultdict +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import regex as re +import torch +from safetensors import safe_open +from safetensors.torch import save_file + +CUBIC_FORMAT = "cubic-pack-quantized" +DEFAULT_SOURCE = Path("__YOUR_PATH__/moonshotai/Kimi-K3") +DEFAULT_OUTPUT = Path("__YOUR_PATH__/Kimi-K3-Cubic-2.5Bit") + +# START-END:BITS@GROUP_SIZE. Kimi-K3 has dense layer 0 and MoE layers 1--92. +# This schedule has a 2.498641304-bit effective width: layers 1--3 use W3 +# G256, layers 4--32 use W3 G512, layers 33--91 use W2 G512, and layer 92 +# uses W4 G512. +MOE_SCHEDULE = "1-3:3@256,4-32:3@512,33-91:2@512,92:4@512" + +# Residual-attention scoring projections are copied in source dtype because a +# 7168-to-1 scorer has negligible storage/compute cost and directly affects the +# softmax that selects residual streams. +LINEAR_SCHEDULE = "" + +DEFAULT_DEVICES = "cuda:0" +DEFAULT_SHARD_SIZE_GIB = 3.0 +DEFAULT_ROW_CHUNK_SIZE = -1 +DEFAULT_TENSOR_BATCH_SIZE = 32 +A8_CARRIER_AWARE = True +QUANT_LOSS_STATS: dict[str, dict[str, Any]] = {} + +NUM_MOE_LAYERS = 92 +NUM_EXPERTS = 896 +HIDDEN_SIZE = 7168 +EXPERT_INPUT_SIZE = 3584 +MOE_INTERMEDIATE_SIZE = 3072 +_E2M1_LEVELS = (0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0) + + +@dataclass(frozen=True) +class MoERule: + start_layer: int + end_layer: int + bits: int + group_size: int + + +@dataclass(frozen=True) +class LinearRule: + layer: int + bits: int + group_size: int + + @property + def prefix(self) -> str: + return f"language_model.model.layers.{self.layer}.mlp_res_proj" + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as file: + for chunk in iter(lambda: file.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _validate_scheme(bits: int, group_size: int) -> None: + if bits not in range(1, 9): + raise ValueError(f"Cubic bits must be in 1--8, got {bits}.") + if group_size <= 0 or group_size * bits % 8: + raise ValueError(f"{bits}-bit group size {group_size} is not byte aligned.") + + +def _parse_moe_schedule(value: str) -> list[MoERule]: + rules = [] + covered = set() + for item in value.split(","): + match = re.fullmatch( + r"(\d+)(?:-(\d+))?:(\d+)@(\d+)", + item.strip(), + ) + if match is None: + raise ValueError("MoE schedule entries must use START-END:BITS@GROUP_SIZE.") + start = int(match.group(1)) + end = int(match.group(2) or start) + bits = int(match.group(3)) + group_size = int(match.group(4)) + _validate_scheme(bits, group_size) + if start > end: + raise ValueError(f"Invalid MoE layer range: {item!r}.") + layers = set(range(start, end + 1)) + if layers & covered: + raise ValueError(f"Overlapping MoE schedule entry: {item!r}.") + covered |= layers + rules.append(MoERule(start, end, bits, group_size)) + expected = set(range(1, NUM_MOE_LAYERS + 1)) + if covered != expected: + missing = sorted(expected - covered) + extra = sorted(covered - expected) + raise ValueError( + f"MoE schedule must cover layers 1--92; missing={missing}, extra={extra}." + ) + return rules + + +def _parse_linear_schedule(value: str) -> list[LinearRule]: + if not value.strip(): + return [] + rules = [] + covered = set() + for item in value.split(","): + match = re.fullmatch(r"(\d+):(\d+)@(\d+)", item.strip()) + if match is None: + raise ValueError("Linear schedule entries must use LAYER:BITS@GROUP_SIZE.") + layer = int(match.group(1)) + bits = int(match.group(2)) + group_size = int(match.group(3)) + _validate_scheme(bits, group_size) + if layer in covered: + raise ValueError(f"Duplicate Linear layer: {layer}.") + covered.add(layer) + rules.append(LinearRule(layer, bits, group_size)) + return rules + + +def _moe_rule_for_key(key: str, rules: list[MoERule]) -> MoERule: + match = re.search(r"\.layers\.(\d+)\.", key) + if match is None: + raise ValueError(f"Cannot determine layer for expert tensor {key}.") + layer = int(match.group(1)) + for rule in rules: + if rule.start_layer <= layer <= rule.end_layer: + return rule + raise ValueError(f"No MoE quantization rule for layer {layer}.") + + +def _cubic_levels( + bits: int, + a: torch.Tensor | float, + b: torch.Tensor | float, + *, + device: torch.device, +) -> torch.Tensor: + a_tensor = torch.as_tensor(a, device=device, dtype=torch.float32) + if bits == 1: + return torch.ones( + (*a_tensor.shape, 1), + device=device, + dtype=torch.float32, + ) + magnitude_max = (1 << (bits - 1)) - 1 + b_tensor = torch.as_tensor(b, device=device, dtype=torch.float32) + t = ( + torch.arange( + magnitude_max + 1, + device=device, + dtype=torch.float32, + ) + / magnitude_max + ) + c = 1 - a_tensor - b_tensor + return t * (a_tensor[..., None] + t * (b_tensor[..., None] + t * c[..., None])) + + +def _carrier_levels(levels: torch.Tensor) -> torch.Tensor: + return torch.round(127 * levels) / 127 + + +def _pack_codes(codes: torch.Tensor, bits: int) -> torch.Tensor: + if bits == 1: + if torch.any((codes != -1) & (codes != 1)): + raise ValueError("1-bit Cubic codes must be -1 or +1.") + else: + magnitude_max = (1 << (bits - 1)) - 1 + if torch.any((codes < -magnitude_max) | (codes > magnitude_max)): + raise ValueError("Cubic codes contain a reserved value.") + + shape = codes.shape + num_values = shape[-1] + flat = codes.reshape(-1, num_values).to(torch.int64) + raw = (flat > 0).to(torch.int64) if bits == 1 else flat & ((1 << bits) - 1) + num_bytes = math.ceil(num_values * bits / 8) + packed = torch.zeros( + flat.shape[0], + num_bytes, + dtype=torch.int64, + device=codes.device, + ) + base = torch.arange(num_values, device=codes.device) * bits + for bit in range(bits): + positions = base + bit + byte_indices = positions // 8 + shifts = positions % 8 + values = ((raw >> bit) & 1) << shifts + packed.scatter_add_( + 1, + byte_indices.expand(flat.shape[0], -1), + values, + ) + return packed.to(torch.uint8).reshape(*shape[:-1], num_bytes) + + +def _decode_mxfp4( + packed: torch.Tensor, + scale: torch.Tensor, + device: torch.device, +) -> torch.Tensor: + packed = packed.to(device=device, non_blocking=True) + scale = scale.to(device=device, non_blocking=True) + lookup = torch.tensor( + _E2M1_LEVELS + tuple(-level for level in _E2M1_LEVELS), + device=device, + dtype=torch.float32, + ) + low = lookup[(packed & 0x0F).to(torch.long)] + high = lookup[(packed >> 4).to(torch.long)] + values = torch.stack((low, high), dim=-1).flatten(-2) + scale_f32 = torch.ldexp( + torch.ones_like(scale, dtype=torch.float32), + scale.to(torch.int32) - 127, + ) + return values * scale_f32.repeat_interleave(32, dim=-1) + + +def _reset_quant_loss_stats() -> None: + QUANT_LOSS_STATS.clear() + + +def _record_quant_loss_stats( + bits: int, + group_size: int, + groups: torch.Tensor, + scale: torch.Tensor, + q: torch.Tensor, +) -> None: + """Accumulate final-format errors without synchronizing the GPU.""" + + key = f"{bits}@{group_size}" + bucket = QUANT_LOSS_STATS.get(key) + if bucket is None: + zero = torch.zeros((), device=groups.device, dtype=torch.float32) + bucket = { + "bits": bits, + "group_size": group_size, + "values": 0, + "signal_sse": zero, + "continuous_sse": zero.clone(), + "carrier_sse": zero.clone(), + "clipped_values": torch.zeros((), device=groups.device, dtype=torch.int64), + } + QUANT_LOSS_STATS[key] = bucket + group_count = int(groups.shape[0]) + value_count = int(groups.numel()) + bucket["values"] += value_count + # Bound diagnostic temporaries independently of the conversion row chunk. + # Eight million FP32 values keep each temporary around 32 MiB and avoid a + # second full expert-batch materialization on small GPUs. + stats_group_chunk = max(1, (8 * 1024 * 1024) // group_size) + for start in range(0, group_count, stats_group_chunk): + stop = min(group_count, start + stats_group_chunk) + group_chunk = groups[start:stop] + scale_chunk = scale[start:stop] + q_chunk = q[start:stop] + reconstructed = group_chunk.sign() * scale_chunk[:, None] * q_chunk + continuous_sse = (group_chunk - reconstructed).square().sum() + if bits <= 2: + # Symmetric W1/W2 levels are already exact on the A8 carrier grid. + carrier_sse = continuous_sse + else: + carrier_q = _carrier_levels(q_chunk) + carrier_reconstructed = ( + group_chunk.sign() * scale_chunk[:, None] * carrier_q + ) + carrier_sse = (group_chunk - carrier_reconstructed).square().sum() + bucket["signal_sse"] = bucket["signal_sse"] + group_chunk.square().sum() + bucket["continuous_sse"] = bucket["continuous_sse"] + continuous_sse + bucket["carrier_sse"] = bucket["carrier_sse"] + carrier_sse + bucket["clipped_values"] = ( + bucket["clipped_values"] + (group_chunk.abs() > scale_chunk[:, None]).sum() + ) + + +def _quant_loss_stats_snapshot() -> dict[str, dict[str, int | float]]: + """Synchronize once per completed source shard and return plain scalars.""" + + result: dict[str, dict[str, int | float]] = {} + for key, bucket in QUANT_LOSS_STATS.items(): + result[key] = { + name: ( + int(value) + if name in ("bits", "group_size", "values") + else float(value.item()) + if isinstance(value, torch.Tensor) + else float(value) + ) + for name, value in bucket.items() + } + return result + + +def _merge_quant_loss_stats( + destination: dict[str, dict[str, int | float]], + source: dict[str, dict[str, int | float]], +) -> None: + for key, incoming in source.items(): + bucket = destination.get(key) + if bucket is None: + destination[key] = dict(incoming) + continue + for name, value in incoming.items(): + if name in ("bits", "group_size"): + if bucket[name] != value: + raise ValueError(f"Inconsistent loss statistic {key} {name}.") + else: + bucket[name] += value + + +def _summarize_quant_loss_bucket( + bucket: dict[str, int | float], + a8_carrier_aware: bool, +) -> dict[str, int | float]: + values = int(bucket["values"]) + signal_sse = max(float(bucket["signal_sse"]), 1.0e-30) + continuous_sse = float(bucket["continuous_sse"]) + carrier_sse = float(bucket["carrier_sse"]) + objective_sse = ( + 0.5 * (continuous_sse + carrier_sse) if a8_carrier_aware else continuous_sse + ) + result: dict[str, int | float] = { + "bits": int(bucket["bits"]), + "loss": math.sqrt(objective_sse / signal_sse), + "clipped_percent": 100.0 * int(bucket["clipped_values"]) / max(values, 1), + } + if a8_carrier_aware: + result["a8_correction_loss"] = math.sqrt(carrier_sse / signal_sse) + if int(bucket["group_size"]) > 0: + result["group_size"] = int(bucket["group_size"]) + return result + + +def _quant_loss_report( + raw_by_scheme: dict[str, dict[str, int | float]], + a8_carrier_aware: bool, +) -> dict[str, Any]: + by_bit_raw: dict[str, dict[str, int | float]] = {} + for bucket in raw_by_scheme.values(): + bit_key = str(int(bucket["bits"])) + aggregate = by_bit_raw.get(bit_key) + if aggregate is None: + aggregate = dict(bucket) + aggregate["group_size"] = -1 + by_bit_raw[bit_key] = aggregate + else: + for name, value in bucket.items(): + if name in ("bits", "group_size"): + continue + aggregate[name] += value + return { + "objective": ( + "mean(continuous MSE, rounded-A8-carrier MSE)" + if a8_carrier_aware + else "continuous MSE" + ), + "normalization": ( + "loss is sqrt(joint SSE / source weight SSE), i.e. NRMSE; " + "the fitting objective itself remains least squares." + ), + "loss_metadata_precision": "FP32 scale and FP16 a/b", + "by_bit": { + key: _summarize_quant_loss_bucket(value, a8_carrier_aware) + for key, value in sorted(by_bit_raw.items(), key=lambda item: int(item[0])) + }, + "by_bit_and_group_size": { + key: _summarize_quant_loss_bucket(value, a8_carrier_aware) + for key, value in sorted( + raw_by_scheme.items(), + key=lambda item: ( + int(item[1]["bits"]), + int(item[1]["group_size"]), + ), + ) + }, + } + + +def _fit_scale( + values: torch.Tensor, + levels: torch.Tensor, + start_scale: torch.Tensor, + iterations: int, +) -> tuple[torch.Tensor, torch.Tensor]: + absolute = values.abs() + carrier_levels = _carrier_levels(levels) if A8_CARRIER_AWARE else None + scale = start_scale.clamp_min(torch.finfo(torch.float32).tiny) + for _ in range(iterations): + candidates = scale[:, None, None] * levels[None, None, :] + distances = (absolute[..., None] - candidates).abs() + if carrier_levels is not None: + carrier_candidates = scale[:, None, None] * carrier_levels[None, None, :] + distances = distances.square() + distances += (absolute[..., None] - carrier_candidates).square() + indices = distances.argmin(dim=-1) + q = levels[indices] + if carrier_levels is None: + numerator = (absolute * q).sum(dim=-1) + denominator = q.square().sum(dim=-1) + else: + carrier_q = carrier_levels[indices] + numerator = (absolute * (q + carrier_q)).sum(dim=-1) + denominator = (q.square() + carrier_q.square()).sum(dim=-1) + updated = numerator / denominator.clamp_min(torch.finfo(torch.float32).tiny) + scale = torch.where(denominator > 0, updated, scale) + q = levels[indices] + reconstructed = values.sign() * scale[:, None] * q + loss = (values - reconstructed).square().sum(dim=-1) + if carrier_levels is not None: + carrier_q = carrier_levels[indices] + carrier_reconstructed = values.sign() * scale[:, None] * carrier_q + loss += (values - carrier_reconstructed).square().sum(dim=-1) + return scale, loss + + +def _quantize_symmetric( + weight: torch.Tensor, + bits: int, + group_size: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + groups = weight.reshape(-1, group_size) + absolute = groups.abs() + if bits == 1: + scale = absolute.mean(dim=-1).clamp_min(torch.finfo(torch.float32).tiny) + codes = torch.where(groups < 0, -1, 1).to(torch.int8) + else: + magnitude_max = (1 << (bits - 1)) - 1 + scale = absolute.amax(dim=-1).clamp_min(torch.finfo(torch.float32).tiny) + for _ in range(8): + magnitudes = torch.round(absolute / scale[:, None] * magnitude_max).clamp_( + 0, magnitude_max + ) + levels = magnitudes / magnitude_max + denominator = levels.square().sum(dim=-1) + updated = (absolute * levels).sum(dim=-1) / denominator.clamp_min( + torch.finfo(torch.float32).tiny + ) + scale = torch.where(denominator > 0, updated, scale) + codes = (groups.sign() * magnitudes).to(torch.int8) + final_q = ( + torch.ones_like(groups) + if bits == 1 + else magnitudes.to(torch.float32) / magnitude_max + ) + _record_quant_loss_stats(bits, group_size, groups, scale, final_q) + codes = codes.reshape(weight.shape) + metadata_shape = ( + *weight.shape[:-1], + weight.shape[-1] // group_size, + ) + scale = scale.reshape(metadata_shape).to(torch.float32) + a = torch.ones( + metadata_shape, + device=weight.device, + dtype=torch.float16, + ) + b = torch.zeros_like(a) + return _pack_codes(codes, bits), scale, a, b + + +def _quantize_curved( + weight: torch.Tensor, + bits: int, + group_size: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + groups = weight.to(torch.float32).reshape(-1, group_size) + calibration_stride = { + 3: 1, + 4: 1, + 5: 2, + 6: 2, + 7: 8, + 8: 8, + }[bits] + calibration = groups[:, ::calibration_stride] + group_amax = groups.abs().amax(dim=-1).clamp_min(torch.finfo(torch.float32).tiny) + best_loss = torch.full_like(group_amax, torch.inf) + best_scale = group_amax.clone() + best_a = torch.ones_like(group_amax) + best_b = torch.zeros_like(group_amax) + candidate_pairs = { + 3: ( + (1.0, 0.0), + (0.75, -0.25), + (1.0, -0.75), + (0.25, 0.25), + (0.5, 0.25), + ), + 4: ( + (0.75, -0.75), + (0.5, -0.25), + (1.0, -0.25), + (0.5, 0.0), + (1.0, 0.0), + ), + 5: ( + (1.0, -0.25), + (1.0, 0.25), + (1.0, 0.0), + (0.75, -0.25), + (0.25, 0.25), + (0.5, -0.25), + (0.75, -0.75), + (0.25, 0.0), + ), + 6: ( + (0.5, -0.75), + (1.25, 0.25), + (1.0, 0.0), + (0.75, -0.25), + (1.25, 0.0), + ), + 7: ( + (0.75, 0.75), + (0.5, -0.75), + (1.25, 0.25), + (1.0, 0.0), + (0.5, -0.25), + ), + 8: ( + (0.5, 0.0), + (0.25, 0.25), + (0.5, -0.75), + (0.75, 0.25), + (1.0, 0.0), + ), + }[bits] + multipliers = { + 3: (0.65, 1.0), + 4: (0.65, 1.0), + 5: (0.65, 0.8, 1.0, 1.15), + 6: (0.65, 0.8, 1.0, 1.15), + 7: (1.0,), + 8: (0.65, 0.8, 1.0, 1.15), + }[bits] + iterations = 8 if bits == 3 else 2 + for a_value, b_value in candidate_pairs: + levels = _cubic_levels( + bits, + a_value, + b_value, + device=weight.device, + ) + for multiplier in multipliers: + scale, loss = _fit_scale( + calibration, + levels, + group_amax * multiplier, + iterations, + ) + improved = loss < best_loss + best_loss = torch.where(improved, loss, best_loss) + best_scale = torch.where(improved, scale, best_scale) + best_a = torch.where(improved, a_value, best_a) + best_b = torch.where(improved, b_value, best_b) + + stored_scale = best_scale.to(torch.float32) + stored_a = best_a.to(torch.float16) + stored_b = best_b.to(torch.float16) + levels = _cubic_levels( + bits, + stored_a, + stored_b, + device=weight.device, + ) + distances = ( + groups.abs()[..., None] - stored_scale[:, None, None] * levels[:, None, :] + ).abs() + if A8_CARRIER_AWARE: + carrier_levels = _carrier_levels(levels) + distances = distances.square() + distances += ( + groups.abs()[..., None] + - stored_scale[:, None, None] * carrier_levels[:, None, :] + ).square() + magnitudes = distances.argmin(dim=-1) + final_q = torch.gather(levels, 1, magnitudes) + _record_quant_loss_stats( + bits, + group_size, + groups, + stored_scale, + final_q, + ) + codes = (groups.sign().to(torch.int64) * magnitudes).reshape(weight.shape) + metadata_shape = ( + *weight.shape[:-1], + weight.shape[-1] // group_size, + ) + return ( + _pack_codes(codes, bits), + stored_scale.reshape(metadata_shape), + stored_a.reshape(metadata_shape), + stored_b.reshape(metadata_shape), + ) + + +def _quantize_weight( + weight: torch.Tensor, + bits: int, + group_size: int, + row_chunk_size: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + if weight.shape[-1] % group_size: + raise ValueError( + f"Weight K={weight.shape[-1]} is not divisible by group size {group_size}." + ) + if bits <= 2: + return _quantize_symmetric(weight, bits, group_size) + if bits == 3: + return _quantize_curved(weight, bits, group_size) + + if row_chunk_size == -1: + row_chunk_size = _automatic_row_chunk_size(weight, bits, group_size) + + packed_chunks = [] + scale_chunks = [] + a_chunks = [] + b_chunks = [] + rows = weight.reshape(-1, weight.shape[-1]) + for chunk in rows.split(row_chunk_size): + packed, scale, a, b = _quantize_curved( + chunk, + bits, + group_size, + ) + packed_chunks.append(packed) + scale_chunks.append(scale) + a_chunks.append(a) + b_chunks.append(b) + prefix = weight.shape[:-1] + return ( + torch.cat(packed_chunks).reshape(*prefix, -1), + torch.cat(scale_chunks).reshape(*prefix, -1), + torch.cat(a_chunks).reshape(*prefix, -1), + torch.cat(b_chunks).reshape(*prefix, -1), + ) + + +def _automatic_row_chunk_size( + weight: torch.Tensor, + bits: int, + group_size: int, +) -> int: + rows = weight.numel() // weight.shape[-1] + if weight.device.type != "cuda": + return min(rows, 32) + free_bytes, _ = torch.accelerator.get_memory_info(weight.device) + level_count = 1 << (bits - 1) + groups_per_row = math.ceil(weight.shape[-1] / group_size) + distance_bytes_per_row = groups_per_row * group_size * level_count * 4 + temporary_factor = 7 if A8_CARRIER_AWARE else 4 + budget = min(int(free_bytes * 0.1), 12 * 1024**3) + estimated = max(1, budget // (distance_bytes_per_row * temporary_factor)) + estimated = min(rows, estimated, 2048) + return estimated // 32 * 32 if estimated >= 32 else estimated + + +def _is_strictly_monotonic(a: float, b: float) -> bool: + c = 1 - a - b + points = [0.0, 1.0] + if c: + vertex = -b / (3 * c) + if 0 < vertex < 1: + points.append(vertex) + return min(a + 2 * b * point + 3 * c * point * point for point in points) > 0 + + +def _linear_candidate_pairs() -> list[tuple[float, float]]: + pairs = [(1.0, 0.0)] + for a in (0.25, 0.5, 0.75, 1.0, 1.25, 1.5): + for b in (-0.75, -0.25, 0.0, 0.25, 0.75): + if (a, b) not in pairs and _is_strictly_monotonic(a, b): + pairs.append((a, b)) + return pairs + + +def _quantize_linear_weight( + weight: torch.Tensor, + bits: int, + group_size: int, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + if weight.shape[-1] % group_size: + raise ValueError( + f"Linear K={weight.shape[-1]} is not divisible by group size {group_size}." + ) + groups = weight.to(torch.float32).reshape(-1, group_size) + group_amax = groups.abs().amax(dim=-1).clamp_min(torch.finfo(torch.float32).tiny) + if bits == 1: + scale = groups.abs().mean(dim=-1).clamp_min(torch.finfo(torch.float32).tiny) + a = torch.ones_like(scale, dtype=torch.float16) + b = torch.zeros_like(a) + codes = torch.where(groups < 0, -1, 1).to(torch.int64) + else: + best_loss = torch.full_like(group_amax, torch.inf) + scale = group_amax.clone() + best_a = torch.ones_like(group_amax) + best_b = torch.zeros_like(group_amax) + for a_value, b_value in _linear_candidate_pairs(): + levels = _cubic_levels( + bits, + a_value, + b_value, + device=weight.device, + ) + for multiplier in (0.65, 0.8, 1.0, 1.15): + candidate_scale, loss = _fit_scale( + groups, + levels, + group_amax * multiplier, + 8, + ) + improved = loss < best_loss + best_loss = torch.where(improved, loss, best_loss) + scale = torch.where(improved, candidate_scale, scale) + best_a = torch.where(improved, a_value, best_a) + best_b = torch.where(improved, b_value, best_b) + scale = scale.to(torch.float32) + a = best_a.to(torch.float16) + b = best_b.to(torch.float16) + levels = _cubic_levels(bits, a, b, device=weight.device) + distances = ( + groups.abs()[..., None] - scale[:, None, None] * levels[:, None, :] + ).abs() + if A8_CARRIER_AWARE: + carrier_levels = _carrier_levels(levels) + distances = distances.square() + distances += ( + groups.abs()[..., None] + - scale[:, None, None] * carrier_levels[:, None, :] + ).square() + magnitudes = distances.argmin(dim=-1) + final_q = torch.gather(levels, 1, magnitudes) + _record_quant_loss_stats( + bits, + group_size, + groups, + scale, + final_q, + ) + codes = groups.sign().to(torch.int64) * magnitudes + if bits == 1: + _record_quant_loss_stats( + bits, + group_size, + groups, + scale, + torch.ones_like(groups), + ) + codes = codes.reshape(weight.shape) + metadata_shape = ( + *weight.shape[:-1], + weight.shape[-1] // group_size, + ) + return ( + _pack_codes(codes, bits), + scale.reshape(metadata_shape).to(torch.float32), + a.reshape(metadata_shape), + b.reshape(metadata_shape), + ) + + +def _quantize_linear( + weight: torch.Tensor, + rule: LinearRule, + device: torch.device, +) -> dict[str, torch.Tensor]: + weight = weight.to(device=device, dtype=torch.float32) + packed, scale, a, b = _quantize_linear_weight( + weight, + rule.bits, + rule.group_size, + ) + return { + f"{rule.prefix}.weight_packed": packed.cpu(), + f"{rule.prefix}.weight_scale": scale.cpu(), + f"{rule.prefix}.weight_a": a.cpu(), + f"{rule.prefix}.weight_b": b.cpu(), + } + + +def _tensor_bytes(tensor: torch.Tensor) -> int: + return tensor.numel() * tensor.element_size() + + +def _format_duration(seconds: float) -> str: + total_seconds = max(0, round(seconds)) + hours, remainder = divmod(total_seconds, 3600) + minutes, seconds = divmod(remainder, 60) + return f"{hours:02d}:{minutes:02d}:{seconds:02d}" + + +def _timed_print( + message: str, + started_at: float, + eta_seconds: float | None, +) -> None: + elapsed = _format_duration(time.monotonic() - started_at) + eta = _format_duration(eta_seconds) if eta_seconds is not None else "calculating" + print(f"[elapsed {elapsed} | ETA {eta}] {message}", flush=True) + + +def _estimate_eta( + started_at: float, + completed: int, + total: int, +) -> float | None: + if completed <= 0: + return None + elapsed = time.monotonic() - started_at + return elapsed * (total - completed) / completed + + +def _estimate_parallel_eta( + completed: int, + total: int, + worker_count: int, + task_durations: list[float], + active_started_at: dict[int, float], +) -> float | None: + if completed < min(worker_count, total): + return None + typical = statistics.median(task_durations[-2 * worker_count :]) + now = time.monotonic() + if any( + now - task_started_at > 2 * typical + for task_started_at in active_started_at.values() + ): + return None + active_progress = [ + min(now - task_started_at, typical) + for task_started_at in active_started_at.values() + ] + remaining_work = max( + 0.0, + (total - completed) * typical - sum(active_progress), + ) + active_tail = max( + (typical - progress for progress in active_progress), + default=0.0, + ) + return max(remaining_work / worker_count, active_tail) + + +class MaterializedShardWriter: + def __init__( + self, + directory: Path, + task_index: int, + max_bytes: int, + progress_callback: Callable[[str], None] | None = None, + ) -> None: + self.directory = directory + self.task_index = task_index + self.max_bytes = max_bytes + self.progress_callback = progress_callback + header_reserve = min(64 * 1024**2, max_bytes // 16) + self.payload_limit = max_bytes - header_reserve + self.part = 0 + self.tensors: dict[str, torch.Tensor] = {} + self.payload_bytes = 0 + self.names: list[str] = [] + + def add(self, tensors: dict[str, torch.Tensor]) -> None: + size = sum(_tensor_bytes(tensor) for tensor in tensors.values()) + if size > self.payload_limit: + names = ", ".join(tensors) + raise ValueError( + f"Tensor group {names} needs {size} bytes, above the " + f"configured shard payload limit {self.payload_limit}." + ) + if self.tensors and self.payload_bytes + size > self.payload_limit: + self.flush() + duplicate = self.tensors.keys() & tensors.keys() + if duplicate: + raise KeyError(f"Duplicate output tensors: {sorted(duplicate)}.") + self.tensors.update(tensors) + self.payload_bytes += size + + def flush(self) -> None: + if not self.tensors: + return + self.part += 1 + name = f"part-task-{self.task_index:06d}-{self.part:06d}.safetensors" + destination = self.directory / name + temporary = destination.with_suffix(".safetensors.tmp") + save_file(self.tensors, temporary) + os.replace(temporary, destination) + actual_size = destination.stat().st_size + if actual_size > self.max_bytes: + raise ValueError( + f"{name} is {actual_size} bytes, above max {self.max_bytes}." + ) + self.names.append(name) + if self.progress_callback is not None: + self.progress_callback(f"wrote {name} ({actual_size / 1024**3:.3f} GiB)") + self.tensors = {} + self.payload_bytes = 0 + + def close(self) -> list[str]: + self.flush() + return self.names + + +def _convert_source_shard( + source_shard: Path, + writer: MaterializedShardWriter, + device: torch.device, + moe_rules: list[MoERule], + linear_rules: dict[str, LinearRule], + row_chunk_size: int, + tensor_batch_size: int, +) -> None: + with safe_open(source_shard, framework="pt", device="cpu") as reader: + keys = list(reader.keys()) + key_set = set(keys) + packed_keys = [ + key + for key in keys + if key.endswith(".weight_packed") + and key.removesuffix(".weight_packed") + ".weight_scale" in key_set + ] + source_scale_keys = { + key.removesuffix(".weight_packed") + ".weight_scale" for key in packed_keys + } + + for key in keys: + if key in packed_keys or key in source_scale_keys: + continue + linear_rule = linear_rules.get(key) + if linear_rule is None: + writer.add({key: reader.get_tensor(key)}) + else: + writer.add( + _quantize_linear( + reader.get_tensor(key), + linear_rule, + device, + ) + ) + + shape_groups: dict[ + tuple[tuple[int, ...], tuple[int, ...], int, int], + list[str], + ] = defaultdict(list) + for key in packed_keys: + base = key.removesuffix(".weight_packed") + rule = _moe_rule_for_key(key, moe_rules) + shape_groups[ + ( + tuple(reader.get_slice(key).get_shape()), + tuple(reader.get_slice(base + ".weight_scale").get_shape()), + rule.bits, + rule.group_size, + ) + ].append(key) + + for (_, _, bits, group_size), grouped_keys in shape_groups.items(): + for start in range(0, len(grouped_keys), tensor_batch_size): + batch_keys = grouped_keys[start : start + tensor_batch_size] + source_packed = torch.stack( + [reader.get_tensor(key) for key in batch_keys] + ) + source_scale = torch.stack( + [ + reader.get_tensor( + key.removesuffix(".weight_packed") + ".weight_scale" + ) + for key in batch_keys + ] + ) + weights = _decode_mxfp4( + source_packed, + source_scale, + device, + ) + packed, scale, a, b = _quantize_weight( + weights, + bits, + group_size, + row_chunk_size, + ) + packed = packed.cpu() + scale = scale.cpu() + a = a.cpu() + b = b.cpu() + for index, key in enumerate(batch_keys): + base = key.removesuffix(".weight_packed") + writer.add( + { + key: packed[index].clone(), + base + ".weight_scale": scale[index].clone(), + base + ".weight_a": a[index].clone(), + base + ".weight_b": b[index].clone(), + } + ) + del ( + source_packed, + source_scale, + weights, + packed, + scale, + a, + b, + ) + + +def _worker( + worker_index: int, + device_name: str, + task_queue: Any, + progress_queue: Any, + source: str, + staging: str, + moe_rules: list[MoERule], + linear_rules: list[LinearRule], + max_shard_bytes: int, + row_chunk_size: int, + tensor_batch_size: int, + a8_carrier_aware: bool, +) -> None: + global A8_CARRIER_AWARE + A8_CARRIER_AWARE = a8_carrier_aware + task_index = -1 + shard_name = "" + try: + torch.set_num_threads(1) + device = torch.device(device_name) + if device.type == "cuda": + torch.accelerator.set_device_index(device.index or 0) + source_path = Path(source) + staging_path = Path(staging) + linear_by_weight = {f"{rule.prefix}.weight": rule for rule in linear_rules} + while True: + task = task_queue.get() + if task is None: + return + task_index, shard_name = task + task_started_at = time.monotonic() + _reset_quant_loss_stats() + progress_queue.put( + ("started", worker_index, device_name, task_index, shard_name) + ) + + def report( + message: str, + current_task_index: int = task_index, + current_shard_name: str = shard_name, + ) -> None: + progress_queue.put( + ( + "log", + worker_index, + device_name, + current_task_index, + current_shard_name, + message, + ) + ) + + writer = MaterializedShardWriter( + staging_path, + task_index, + max_shard_bytes, + report, + ) + _convert_source_shard( + source_path / shard_name, + writer, + device, + moe_rules, + linear_by_weight, + row_chunk_size, + tensor_batch_size, + ) + if device.type == "cuda": + torch.accelerator.empty_cache() + names = writer.close() + progress_queue.put( + ( + "completed", + worker_index, + device_name, + task_index, + shard_name, + time.monotonic() - task_started_at, + len(names), + _quant_loss_stats_snapshot(), + ) + ) + except BaseException: + progress_queue.put( + ( + "error", + worker_index, + device_name, + task_index, + shard_name, + traceback.format_exc(), + ) + ) + + +def _run_workers( + devices: list[str], + source_shards: list[str], + source: Path, + staging: Path, + moe_rules: list[MoERule], + linear_rules: list[LinearRule], + max_shard_bytes: int, + row_chunk_size: int, + tensor_batch_size: int, + a8_carrier_aware: bool, + started_at: float, +) -> dict[str, dict[str, int | float]]: + # Workers pull the next source shard only after finishing their current + # shard. Output part names use the source-shard index, so scheduling does + # not change tensor-to-part assignment. + context = mp.get_context("spawn") + task_queue = context.Queue() + progress_queue = context.Queue() + for task in enumerate(source_shards): + task_queue.put(task) + for _ in devices: + task_queue.put(None) + + arguments = [ + ( + index, + device, + task_queue, + progress_queue, + str(source), + str(staging), + moe_rules, + linear_rules, + max_shard_bytes, + row_chunk_size, + tensor_batch_size, + a8_carrier_aware, + ) + for index, device in enumerate(devices) + ] + processes = [ + context.Process(target=_worker, args=worker_args) for worker_args in arguments + ] + for process in processes: + process.start() + + total = len(source_shards) + completed = 0 + task_durations: list[float] = [] + active_started_at: dict[int, float] = {} + pending = set(range(len(processes))) + empty_after_exit = 0 + quant_loss_stats: dict[str, dict[str, int | float]] = {} + try: + while pending or completed < total: + try: + event = progress_queue.get(timeout=1) + except queue.Empty: + event = None + if event is None and not pending: + empty_after_exit += 1 + if empty_after_exit >= 5: + raise RuntimeError( + f"Workers completed only {completed}/{total} source shards." + ) + else: + empty_after_exit = 0 + if event is not None: + kind = event[0] + if kind == "started": + _, worker, device, task_index, shard_name = event + active_started_at[worker] = time.monotonic() + eta = _estimate_parallel_eta( + completed, + total, + len(devices), + task_durations, + active_started_at, + ) + _timed_print( + f"[worker {worker} {device}] started " + f"task {task_index + 1}/{total}: {shard_name}", + started_at, + eta, + ) + elif kind == "log": + _, worker, device, task_index, shard_name, message = event + eta = _estimate_parallel_eta( + completed, + total, + len(devices), + task_durations, + active_started_at, + ) + _timed_print( + f"[worker {worker} {device}] task " + f"{task_index + 1}/{total} {shard_name}: {message}", + started_at, + eta, + ) + elif kind == "completed": + ( + _, + worker, + device, + task_index, + shard_name, + task_seconds, + part_count, + task_loss_stats, + ) = event + _merge_quant_loss_stats(quant_loss_stats, task_loss_stats) + completed += 1 + active_started_at.pop(worker, None) + task_durations.append(task_seconds) + eta = _estimate_parallel_eta( + completed, + total, + len(devices), + task_durations, + active_started_at, + ) + _timed_print( + f"[worker {worker} {device}] completed " + f"{completed}/{total}: {shard_name}; " + f"task time {_format_duration(task_seconds)}, " + f"{part_count} output parts", + started_at, + eta, + ) + elif kind == "error": + _, worker, device, task_index, shard_name, details = event + raise RuntimeError( + f"Worker {worker} ({device}) failed on task " + f"{task_index + 1} ({shard_name}):\n{details}" + ) + else: + raise RuntimeError(f"Unknown worker progress event: {kind!r}.") + + for index in tuple(pending): + returncode = processes[index].exitcode + if returncode is None: + continue + pending.remove(index) + if returncode: + raise RuntimeError( + f"Quantization worker {index} exited with code {returncode}." + ) + except BaseException: + for index in pending: + processes[index].terminate() + raise + finally: + for process in processes: + process.join() + task_queue.close() + progress_queue.close() + return quant_loss_stats + + +def _weights_config(bits: int, group_size: int) -> dict: + return { + "num_bits": bits, + "group_size": group_size, + "strategy": "group", + "symmetric": True, + "dynamic": False, + "scale_dtype": "torch.float32", + "param_dtype": "torch.float16", + "reserved_code": "binary" if bits == 1 else "zero", + "packing": "little-endian-bitstream", + } + + +def _layer_target(start: int, end: int) -> str: + pattern = "|".join(str(layer) for layer in range(start, end + 1)) + return ( + rf"re:.*\.layers\.(?:{pattern})\." + r"block_sparse_moe\.experts" + ) + + +def _effective_bits( + moe_rules: list[MoERule], + linear_rules: list[LinearRule], +) -> tuple[float, float, float]: + rule_by_layer = { + layer: rule + for rule in moe_rules + for layer in range(rule.start_layer, rule.end_layer + 1) + } + expert_bit_sum = sum( + rule_by_layer[layer].bits + 64 / rule_by_layer[layer].group_size + for layer in range(1, NUM_MOE_LAYERS + 1) + ) + payload = ( + sum(rule_by_layer[layer].bits for layer in range(1, NUM_MOE_LAYERS + 1)) + / NUM_MOE_LAYERS + ) + expert_effective = expert_bit_sum / NUM_MOE_LAYERS + expert_values_per_layer = ( + NUM_EXPERTS * 3 * EXPERT_INPUT_SIZE * MOE_INTERMEDIATE_SIZE + ) + linear_bit_sum = sum(rule.bits + 64 / rule.group_size for rule in linear_rules) + converted_effective = ( + expert_values_per_layer * expert_bit_sum + HIDDEN_SIZE * linear_bit_sum + ) / (expert_values_per_layer * NUM_MOE_LAYERS + HIDDEN_SIZE * len(linear_rules)) + return payload, expert_effective, converted_effective + + +def _quantization_config( + source_config: dict, + moe_rules: list[MoERule], + linear_rules: list[LinearRule], +) -> dict: + config_groups = {} + for rule in moe_rules: + name = f"moe_layers_{rule.start_layer}_{rule.end_layer}" + config_groups[name] = { + "targets": [_layer_target(rule.start_layer, rule.end_layer)], + "input_activations": None, + "output_activations": None, + "weights": _weights_config(rule.bits, rule.group_size), + } + for rule in linear_rules: + config_groups[f"linear_layer_{rule.layer}_{rule.bits}bit"] = { + "targets": [rule.prefix], + "input_activations": None, + "output_activations": None, + "weights": _weights_config(rule.bits, rule.group_size), + } + payload, expert_effective, converted_effective = _effective_bits( + moe_rules, + linear_rules, + ) + group_overrides = { + str(layer): rule.group_size + for rule in moe_rules + for layer in range(rule.start_layer, rule.end_layer + 1) + } + return { + "quant_method": "cubic", + "format": CUBIC_FORMAT, + "quantization_status": "compressed", + "config_groups": config_groups, + "ignore": source_config.get("ignore", []), + "runtime_weight_storage": "native_packed_bitstream", + "layer_bit_schedule": [ + { + "start_layer": rule.start_layer, + "end_layer": rule.end_layer, + "num_bits": rule.bits, + "group_size": rule.group_size, + } + for rule in moe_rules + ], + "layer_group_size_overrides": group_overrides, + "tensor_bit_overrides": [ + { + "target": rule.prefix, + "num_bits": rule.bits, + "group_size": rule.group_size, + } + for rule in linear_rules + ], + "expert_payload_bits": payload, + "expert_effective_bits": expert_effective, + "converted_tensor_effective_bits": converted_effective, + } + + +def _copy_model_assets(source: Path, staging: Path) -> None: + excluded = { + "config.json", + "model.safetensors.index.json", + "cubic_quantization_manifest.json", + "cubic_quantization_audit.json", + "cubic_quantization_report.json", + } + for path in source.iterdir(): + if ( + path.name in excluded + or path.name.startswith("model-") + and path.suffix == ".safetensors" + ): + continue + destination = staging / path.name + if path.is_dir(): + shutil.copytree(path, destination, symlinks=False) + else: + shutil.copy2(path, destination, follow_symlinks=True) + + +def _finalize_shards( + staging: Path, + progress_callback: Callable[[int, int, str], None] | None = None, +) -> tuple[dict[str, str], int, list[str]]: + parts = sorted(staging.glob("part-task-*.safetensors")) + if not parts: + raise RuntimeError("No output safetensors were produced.") + total = len(parts) + width = max(5, len(str(total))) + final_names = [ + f"model-{index:0{width}d}-of-{total:0{width}d}.safetensors" + for index in range(1, total + 1) + ] + for part, name in zip(parts, final_names, strict=True): + os.replace(part, staging / name) + + weight_map = {} + total_size = 0 + for position, name in enumerate(final_names, start=1): + path = staging / name + total_size += path.stat().st_size + with safe_open(path, framework="pt", device="cpu") as reader: + keys = reader.keys() + for key in keys: + if key in weight_map: + raise KeyError(f"Duplicate tensor in output: {key}.") + weight_map[key] = name + if progress_callback is not None: + progress_callback(position, total, name) + return weight_map, total_size, final_names + + +def _expected_output_keys( + source_weight_map: dict[str, str], + linear_rules: list[LinearRule], +) -> set[str]: + linear_weights = {f"{rule.prefix}.weight": rule for rule in linear_rules} + packed_keys = { + key + for key in source_weight_map + if key.endswith(".weight_packed") + and key.removesuffix(".weight_packed") + ".weight_scale" in source_weight_map + } + source_scales = { + key.removesuffix(".weight_packed") + ".weight_scale" for key in packed_keys + } + expected = set() + for key in source_weight_map: + if key in source_scales: + continue + if key in packed_keys: + base = key.removesuffix(".weight_packed") + expected.update( + { + key, + base + ".weight_scale", + base + ".weight_a", + base + ".weight_b", + } + ) + elif key in linear_weights: + base = key.removesuffix(".weight") + expected.update( + { + base + ".weight_packed", + base + ".weight_scale", + base + ".weight_a", + base + ".weight_b", + } + ) + else: + expected.add(key) + return expected + + +def _audit( + model: Path, + max_shard_bytes: int, + expected_keys: set[str], + expected_widths: set[int], + progress_callback: Callable[[int, int, str], None] | None = None, +) -> dict: + if any(path.is_symlink() for path in model.rglob("*")): + raise ValueError("Output contains a symbolic link.") + config = json.loads((model / "config.json").read_text()) + text_config = config.get("text_config", config) + quantization = text_config["quantization_config"] + if quantization.get("expert_placement") is not None: + raise ValueError("Output unexpectedly contains expert_placement.") + if ( + quantization.get("quant_method") != "cubic" + or quantization.get("format") != CUBIC_FORMAT + ): + raise ValueError("Output does not declare the Cubic format.") + if quantization["converted_tensor_effective_bits"] > 2.5: + raise ValueError("Converted effective width exceeds 2.5 bits.") + widths = { + group["weights"]["num_bits"] for group in quantization["config_groups"].values() + } + if widths != expected_widths: + raise ValueError( + f"Expected widths {sorted(expected_widths)}, got {sorted(widths)}." + ) + + index = json.loads((model / "model.safetensors.index.json").read_text()) + weight_map = index["weight_map"] + if set(weight_map) != expected_keys: + missing = expected_keys - weight_map.keys() + extra = weight_map.keys() - expected_keys + raise ValueError( + f"Output key mismatch: missing={len(missing)}, extra={len(extra)}." + ) + shard_names = sorted(set(weight_map.values())) + actual_map = {} + dtype_counts = Counter() + total_size = 0 + for position, shard_name in enumerate(shard_names, start=1): + shard = model / shard_name + if shard.is_symlink() or not shard.is_file(): + raise FileNotFoundError(f"Invalid shard: {shard}.") + size = shard.stat().st_size + if size > max_shard_bytes: + raise ValueError(f"{shard_name} exceeds configured shard size.") + total_size += size + with safe_open(shard, framework="pt", device="cpu") as reader: + keys = reader.keys() + for key in keys: + if key in actual_map: + raise KeyError(f"Duplicate tensor: {key}.") + actual_map[key] = shard_name + dtype = reader.get_slice(key).get_dtype() + dtype_counts[str(dtype)] += 1 + if key.endswith(".weight_scale") and dtype != "F32": + raise ValueError(f"Invalid Cubic scale {key}.") + if key.endswith((".weight_a", ".weight_b")) and dtype != "F16": + raise ValueError(f"Invalid Cubic parameter {key}.") + if progress_callback is not None: + progress_callback(position, len(shard_names), shard_name) + if actual_map != weight_map: + raise ValueError("Safetensors headers do not match weight_map.") + if total_size != index["metadata"]["total_size"]: + raise ValueError("Index total_size does not match materialized files.") + return { + "checkpoint": str(model), + "shards": len(shard_names), + "tensors": len(weight_map), + "total_size": total_size, + "max_shard_bytes": max_shard_bytes, + "widths_present": sorted(widths), + "converted_tensor_effective_bits": quantization[ + "converted_tensor_effective_bits" + ], + "dtype_counts": dict(dtype_counts), + } + + +def _preflight( + source: Path, + output: Path, + staging: Path, + devices: list[str], + moe_rules: list[MoERule], + linear_rules: list[LinearRule], +) -> tuple[dict, dict]: + if output.exists(): + raise FileExistsError(f"Output already exists: {output}") + if staging.exists(): + raise FileExistsError( + f"Incomplete output exists: {staging}. " + "Inspect and remove it before retrying." + ) + for name in ("config.json", "model.safetensors.index.json"): + if not (source / name).is_file(): + raise FileNotFoundError(f"Missing source file: {source / name}") + if not devices: + raise ValueError("At least one conversion device is required.") + if len(set(devices)) != len(devices): + raise ValueError("Conversion devices must not contain duplicates.") + + config = json.loads((source / "config.json").read_text()) + index = json.loads((source / "model.safetensors.index.json").read_text()) + source_shards = set(index["weight_map"].values()) + missing_shards = [ + shard for shard in sorted(source_shards) if not (source / shard).is_file() + ] + if missing_shards: + raise FileNotFoundError(f"Missing {len(missing_shards)} source shards.") + linear_weights = {f"{rule.prefix}.weight" for rule in linear_rules} + missing_linear = linear_weights - index["weight_map"].keys() + if missing_linear: + raise KeyError( + f"Linear schedule targets are missing: {sorted(missing_linear)}." + ) + payload, expert_effective, converted_effective = _effective_bits( + moe_rules, + linear_rules, + ) + if converted_effective > 2.5: + raise ValueError( + f"Schedule effective width is {converted_effective}, above 2.5." + ) + return config, index + + +def quantize(args: argparse.Namespace) -> None: + started_at = time.monotonic() + source = args.source.resolve() + output = args.output.resolve() + staging = output.with_name(output.name + ".incomplete") + moe_rules = _parse_moe_schedule(args.moe_schedule) + linear_rules = _parse_linear_schedule(args.linear_schedule) + devices = [device.strip() for device in args.devices.split(",") if device.strip()] + max_shard_bytes = int(args.shard_size_gib * 1024**3) + if max_shard_bytes <= 0: + raise ValueError("--shard-size-gib must be positive.") + config, source_index = _preflight( + source, + output, + staging, + devices, + moe_rules, + linear_rules, + ) + payload, expert_effective, converted_effective = _effective_bits( + moe_rules, + linear_rules, + ) + plan = { + "source": str(source), + "output": str(output), + "temporary_output": str(staging), + "moe_schedule": args.moe_schedule, + "linear_schedule": args.linear_schedule, + "devices": devices, + "shard_size_gib": args.shard_size_gib, + "a8_carrier_aware": args.a8_carrier_aware, + "a8_correction_default": "enabled", + "fitting_objective": "groupwise-least-squares", + "reported_loss": "NRMSE = sqrt(joint SSE / source weight SSE)", + "row_chunk_size": args.row_chunk_size, + "source_shards": len(set(source_index["weight_map"].values())), + "expert_payload_bits": payload, + "expert_effective_bits": expert_effective, + "converted_tensor_effective_bits": converted_effective, + "worker_scheduling": "dynamic_source_shard_queue", + "output_partitioning": "deterministic_source_shard_index", + } + _timed_print( + json.dumps(plan, ensure_ascii=False, indent=2), + started_at, + None, + ) + if args.plan: + total_elapsed = time.monotonic() - started_at + _timed_print( + f"Plan validation completed. Total elapsed: " + f"{_format_duration(total_elapsed)}.", + started_at, + 0, + ) + return + + staging.mkdir(parents=True) + source_shards = sorted(set(source_index["weight_map"].values())) + raw_quant_loss_stats = _run_workers( + devices, + source_shards, + source, + staging, + moe_rules, + linear_rules, + max_shard_bytes, + args.row_chunk_size, + args.tensor_batch_size, + args.a8_carrier_aware, + started_at, + ) + loss_statistics = _quant_loss_report( + raw_quant_loss_stats, + args.a8_carrier_aware, + ) + _timed_print( + "All source shards converted; finalizing output shards.", + started_at, + None, + ) + finalize_started_at = time.monotonic() + + def report_finalize(position: int, total: int, name: str) -> None: + eta = _estimate_eta(finalize_started_at, position, total) + _timed_print( + f"[finalize] indexed {position}/{total}: {name}", + started_at, + eta, + ) + + weight_map, total_size, shard_names = _finalize_shards( + staging, + report_finalize, + ) + expected_keys = _expected_output_keys( + source_index["weight_map"], + linear_rules, + ) + if set(weight_map) != expected_keys: + raise ValueError("Converted keys do not match the source model.") + + text_config = config.get("text_config", config) + source_quantization = text_config.get("quantization_config", {}) + text_config["quantization_config"] = _quantization_config( + source_quantization, + moe_rules, + linear_rules, + ) + index = { + "metadata": {"total_size": total_size}, + "weight_map": weight_map, + } + (staging / "config.json").write_text( + json.dumps(config, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + (staging / "model.safetensors.index.json").write_text( + json.dumps(index, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + _copy_model_assets(source, staging) + + manifest = { + **plan, + "script": str(Path(__file__).resolve()), + "script_sha256": _sha256(Path(__file__).resolve()), + "source_config_sha256": _sha256(source / "config.json"), + "source_index_sha256": _sha256(source / "model.safetensors.index.json"), + "output_shards": len(shard_names), + "output_total_bytes": total_size, + "combined_report": "cubic_quantization_report.json", + } + _timed_print( + "Output metadata written; auditing materialized checkpoint.", + started_at, + None, + ) + audit_started_at = time.monotonic() + + def report_audit(position: int, total: int, name: str) -> None: + eta = _estimate_eta(audit_started_at, position, total) + _timed_print( + f"[audit] checked {position}/{total}: {name}", + started_at, + eta, + ) + + audit = _audit( + staging, + max_shard_bytes, + expected_keys, + { + *(rule.bits for rule in moe_rules), + *(rule.bits for rule in linear_rules), + }, + report_audit, + ) + combined_report = { + "manifest": manifest, + "loss_statistics": loss_statistics, + "audit": audit, + } + (staging / "cubic_quantization_report.json").write_text( + json.dumps(combined_report, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + os.replace(staging, output) + _timed_print( + json.dumps( + { + "loss_by_bit": loss_statistics["by_bit"], + "loss_by_bit_and_group_size": loss_statistics["by_bit_and_group_size"], + "audit": audit, + "combined_report": str(output / "cubic_quantization_report.json"), + }, + ensure_ascii=False, + indent=2, + ), + started_at, + 0, + ) + _timed_print( + f"Quantized checkpoint written to {output}.", + started_at, + 0, + ) + total_elapsed = time.monotonic() - started_at + _timed_print( + f"Total elapsed: {_format_duration(total_elapsed)}.", + started_at, + 0, + ) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--source", + type=Path, + default=DEFAULT_SOURCE, + help="Source Kimi-K3 MXFP4 checkpoint.", + ) + parser.add_argument( + "--output", + type=Path, + default=DEFAULT_OUTPUT, + help="New materialized Cubic checkpoint; must not exist.", + ) + parser.add_argument( + "--moe-schedule", + default=MOE_SCHEDULE, + help="MoE rules formatted START-END:BITS@GROUP_SIZE.", + ) + parser.add_argument( + "--linear-schedule", + default=LINEAR_SCHEDULE, + help=( + "Optional mlp_res_proj rules formatted LAYER:BITS@GROUP_SIZE; " + "disabled by default because these tensors score residual streams." + ), + ) + parser.add_argument( + "--devices", + default=DEFAULT_DEVICES, + help=( + "Comma-separated conversion-only devices, one worker per device; " + "independent of the inference topology." + ), + ) + parser.add_argument( + "--shard-size-gib", + type=float, + default=DEFAULT_SHARD_SIZE_GIB, + help="Maximum size of each materialized safetensors shard.", + ) + parser.add_argument( + "--row-chunk-size", + type=int, + default=DEFAULT_ROW_CHUNK_SIZE, + help=( + "Rows per high-bit calibration chunk; -1 selects a memory-aware " + "value, while a positive value forces an exact chunk size." + ), + ) + parser.add_argument( + "--tensor-batch-size", + type=int, + default=DEFAULT_TENSOR_BATCH_SIZE, + help="MXFP4 expert tensors converted together on each GPU.", + ) + a8_group = parser.add_mutually_exclusive_group() + a8_group.add_argument( + "--a8-carrier-aware", + dest="a8_carrier_aware", + action="store_true", + default=True, + help=( + "Jointly fit continuous Cubic and round(127*q)/127 for Dynamic A8; " + "enabled by default." + ), + ) + a8_group.add_argument( + "--disable-a8-correction", + dest="a8_carrier_aware", + action="store_false", + help=( + "Disable round(127*q)/127 carrier correction and fit only the " + "continuous Cubic reconstruction." + ), + ) + parser.add_argument( + "--plan", + action="store_true", + help="Validate and print the conversion plan without writing.", + ) + args = parser.parse_args() + if args.row_chunk_size == 0 or args.row_chunk_size < -1: + parser.error("--row-chunk-size must be -1 or a positive integer.") + if args.tensor_batch_size <= 0: + parser.error("--tensor-batch-size must be positive.") + return args + + +if __name__ == "__main__": + quantize(parse_args()) diff --git a/tiktoken.model b/tiktoken.model new file mode 100644 index 0000000000000000000000000000000000000000..b4149a6e17a01b6442187f39890f89bc2fe8d309 --- /dev/null +++ b/tiktoken.model @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b6c497a7469b33ced9c38afb1ad6e47f03f5e5dc05f15930799210ec050c5103 +size 2795286 diff --git a/tokenization_kimi.py b/tokenization_kimi.py new file mode 100644 index 0000000000000000000000000000000000000000..5df9d462b43295245084fbab46fd956020d4b84e --- /dev/null +++ b/tokenization_kimi.py @@ -0,0 +1,408 @@ +import os +from logging import getLogger +from pathlib import Path +from shutil import copyfile +from typing import Dict, Iterator, List, Optional, Tuple, Union, cast + +import tiktoken +from tiktoken.load import load_tiktoken_bpe +from tokenizers import AddedToken +from transformers.convert_slow_tokenizer import bytes_to_unicode +from transformers.tokenization_utils import PreTrainedTokenizer + +try: + from .encoding_k3 import build_chat_segments, is_batched_conversation +except ImportError: # pragma: no cover - supports direct file execution/import. + from encoding_k3 import build_chat_segments, is_batched_conversation + +logger = getLogger(__name__) +VOCAB_FILES_NAMES = {"vocab_file": "tiktoken.model"} + + +class TikTokenTokenizer(PreTrainedTokenizer): + """ + Tokenizing and encoding/decoding text using the Tiktoken tokenizer. See megatron/tokenizer/tiktoken_tokenizer.py. + + This tokenizer inherits from [`PreTrainedTokenizer`] which contains most of the main methods. Users should refer to + this superclass for more information regarding those methods. + + Args: + vocab_file (`str`): + The path to the Tiktoken model file. + bos_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|begin_of_text|>",`): + The beginning of sequence token that was used during pretraining. Can be used a sequence classifier token. + eos_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|end_of_text|>"`): + The end of sequence token. + unk_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|reserved_special_token_249|>"`): + The unknown token. A token that is not in the vocabulary cannot be converted to an ID and is set to be this + token instead. The second to last item in special_tokens. + pad_token (`str` or `tokenizers.AddedToken`, *optional*, defaults to `"<|reserved_special_token_250|>"`): + The token used for padding, for example when batching sequences of different lengths. + additional_special_tokens (list of `str`, *optional*): + A tuple or a list of additional tokens, which will be marked as `special`, meaning that they will be + skipped when decoding if `skip_special_tokens` is set to `True`. + """ + + vocab_files_names = VOCAB_FILES_NAMES + + model_input_names = ["input_ids", "attention_mask"] + + special_tokens: Dict[str, int] + + num_reserved_special_tokens = 256 + + pat_str = "|".join([ + r"""[\p{Han}]+""", + r"""[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+(?i:'s|'t|'re|'ve|'m|'ll|'d)?""", + r"""[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]+[\p{Ll}\p{Lm}\p{Lo}\p{M}&&[^\p{Han}]]*(?i:'s|'t|'re|'ve|'m|'ll|'d)?""", + r"""\p{N}{1,3}""", + r""" ?[^\s\p{L}\p{N}]+[\r\n]*""", + r"""\s*[\r\n]+""", + r"""\s+(?!\S)""", + r"""\s+""", + ]) + + def __init__( + self, + vocab_file, + bos_token: Union[str, AddedToken] = "[BOS]", + eos_token: Union[str, AddedToken] = "[EOS]", + unk_token: Union[str, AddedToken, None] = None, + pad_token: Union[str, AddedToken, None] = None, + additional_special_tokens: List[str] = None, + added_tokens_decoder: Optional[dict] = None, + **kwargs, + ): + assert os.path.isfile(vocab_file), vocab_file + + if additional_special_tokens is None: + additional_special_tokens = [ + "<|im_end|>", + "<|im_user|>", + "<|im_assistant|>", + "<|start_header_id|>", + "<|end_header_id|>", + "[EOT]", + "<|im_system|>", + "<|im_middle|>", + ] + + if added_tokens_decoder: + special_tokens_mapping = { + i: added_tokens_decoder[i].content + for i in added_tokens_decoder + } + else: + special_tokens_mapping = {} + + self.vocab_file = vocab_file + mergeable_ranks = load_tiktoken_bpe(vocab_file) + num_base_tokens = len(mergeable_ranks) + self.special_tokens = { + special_tokens_mapping.get(i, f"<|reserved_token_{i}|>"): i + for i in range(num_base_tokens, num_base_tokens + + self.num_reserved_special_tokens) + } + + self.model = tiktoken.Encoding( + name=Path(vocab_file).name, + pat_str=self.pat_str, + mergeable_ranks=mergeable_ranks, + special_tokens=self.special_tokens, + ) + logger.info(f"Reloaded tiktoken model from {vocab_file}") + + self.n_words: int = self.model.n_vocab + # BOS / EOS token IDs + self.bos_id: int = self.special_tokens[str(bos_token)] + self.eos_id: int = self.special_tokens[str(eos_token)] + logger.info( + f"#words: {self.n_words} - BOS ID: {self.bos_id} - EOS ID: {self.eos_id}" + ) + + self.pad_id: int = self.special_tokens[str(pad_token)] + self.unk_id: int = self.special_tokens[str(unk_token)] + + self.byte_encoder = bytes_to_unicode() + self.byte_decoder = {v: k for k, v in self.byte_encoder.items()} + + self.decoder = {} + for i in range(self.n_words): + # Taken from https://gist.github.com/xenova/a452a6474428de0182b17605a98631ee + decoding = ''.join([ + self.byte_encoder[ord(char)] for char in + self.model.decode_single_token_bytes(i).decode('latin-1') + ]) + self.decoder[i] = decoding + + self.encoder = {} + for i in range(self.n_words): + if i in self.decoder: + self.encoder[self.decoder[i]] = i + + super().__init__( + bos_token=bos_token, + eos_token=eos_token, + unk_token=unk_token, + pad_token=pad_token, + additional_special_tokens=additional_special_tokens, + added_tokens_decoder=added_tokens_decoder, + **kwargs, + ) + self.all_special_ids_set = set(self.all_special_ids) + + def _encode_text_piece(self, text: str, + allow_special_tokens: bool = True) -> List[int]: + # The tiktoken tokenizer can handle <=400k chars without + # pyo3_runtime.PanicException. + TIKTOKEN_MAX_ENCODE_CHARS = 400_000 + + # https://github.com/openai/tiktoken/issues/195 + # Here we iterate over subsequences and split if we exceed the limit + # of max consecutive non-whitespace or whitespace characters. + MAX_NO_WHITESPACES_CHARS = 25_000 + + t: List[int] = [] + for i in range(0, len(text), TIKTOKEN_MAX_ENCODE_CHARS): + for substr in self._split_whitespaces_or_nonwhitespaces( + text[i:i + TIKTOKEN_MAX_ENCODE_CHARS], + MAX_NO_WHITESPACES_CHARS, + ): + if allow_special_tokens: + t.extend( + # structural markers: encode <|...|> as their special token IDs + self.model.encode( + substr, + allowed_special="all", + )) + else: + t.extend( + # user/tool text: encode any <|...|> as ordinary BPE tokens (never as control tokens) + self.model.encode( + substr, + disallowed_special=(), + )) + + return t + + def encode(self, + text: str, + allow_special_tokens: bool = True, + **kwargs) -> List[int]: + """ + Encodes a string into a list of token IDs. + + Args: + text (str): The input string to be encoded. + + Returns: + list[int]: A list of token IDs. + """ + # If there are other args, we should call super().encode because there are a lot of code + # to handle those args. supper().encode finally will call _tokenize and _convert_token_to_id. + # NOTE: our encode method is not compatible with the super().encode method, + # e.g. split_special_tokens' default is True in our encode method. + if len(kwargs) > 0: + logger.warning(f"Calling super().encode with {kwargs}") + return super().encode(text, **kwargs) + + assert type(text) is str + return self._encode_text_piece(text, + allow_special_tokens=allow_special_tokens) + + def decode(self, token_ids: Union[int, List[int]], **kwargs) -> str: + """ + Decodes a list of token IDs into a string. + + Args: + token_ids (List[int]): The list of token IDs to be decoded. + + Returns: + str: The decoded string. + """ + # If there are other args, we should call super().decode because there are a lot of code + # to handle those args. supper().encode finally will call convert_tokens_to_string and _convert_id_to_token. + if len(kwargs) > 0: + return super().decode(token_ids, **kwargs) + + if type(token_ids) is int: + token_ids = [token_ids] + + return self.model.decode(cast(List[int], token_ids)) + + @staticmethod + def _split_whitespaces_or_nonwhitespaces( + s: str, max_consecutive_slice_len: int) -> Iterator[str]: + """ + Splits the string `s` so that each substring contains no more than `max_consecutive_slice_len` + consecutive whitespaces or consecutive non-whitespaces. + """ + current_slice_len = 0 + current_slice_is_space = s[0].isspace() if len(s) > 0 else False + slice_start = 0 + + for i in range(len(s)): + is_now_space = s[i].isspace() + + if current_slice_is_space ^ is_now_space: + current_slice_len = 1 + current_slice_is_space = is_now_space + else: + current_slice_len += 1 + if current_slice_len > max_consecutive_slice_len: + yield s[slice_start:i] + slice_start = i + current_slice_len = 1 + yield s[slice_start:] + + def _encode_chat_segments(self, segments) -> List[int]: + token_ids: List[int] = [] + for segment in segments: + token_ids.extend( + self._encode_text_piece( + segment.text, + allow_special_tokens=segment.allow_special, + )) + return token_ids + + @staticmethod + def _truncate(ids: List[int], + truncation: bool = False, + max_length: Optional[int] = None) -> List[int]: + if truncation and max_length is not None: + return ids[:max_length] + return ids + + def _format_chat_token_output(self, + encoded_inputs: List[List[int]], + *, + is_batched: bool, + padding=False, + truncation: bool = False, + max_length: Optional[int] = None, + return_tensors=None, + return_dict: bool = False): + encoded_inputs = [ + self._truncate(ids, truncation=truncation, max_length=max_length) + for ids in encoded_inputs + ] + + needs_batch_encoding = ( + is_batched or padding or return_tensors is not None or return_dict) + if not needs_batch_encoding: + return encoded_inputs[0] + + features = [{ + "input_ids": ids, + "attention_mask": [1] * len(ids) + } for ids in encoded_inputs] + batch = self.pad(features, + padding=padding, + max_length=max_length if padding else None, + return_attention_mask=True, + return_tensors=return_tensors) + + if return_dict: + return batch + if is_batched: + return batch["input_ids"] + return batch["input_ids"][0] if return_tensors is None else batch[ + "input_ids"] + + """ ----- Below are the abstract methods required by PreTrainedTokenizer ----- """ + + @property + def vocab_size(self) -> int: + return self.n_words + + def get_vocab(self) -> Dict[str, int]: + return self.encoder + + def _tokenize(self, text: str, **kwargs) -> List[str]: + return [self.decoder[t] for t in self.encode(text)] + + def _convert_token_to_id(self, token: str) -> int: + return self.encoder.get(token, self.unk_id) + + def _convert_id_to_token(self, index: int) -> str: + return self.decoder.get(index) + + @staticmethod + def clean_up_tokenization(out_string: str) -> str: + return out_string + + def convert_tokens_to_string(self, tokens: List[str]) -> str: + text = ''.join(tokens) + text = bytearray([self.byte_decoder[c] + for c in text]).decode('utf-8', 'replace') + return text + + def save_vocabulary(self, + save_directory: str, + filename_prefix: Optional[str] = None) -> Tuple[str]: + if not os.path.isdir(save_directory): + raise ValueError( + f"vocabulary path ({save_directory}) should be a directory") + out_vocab_file = os.path.join( + save_directory, + (filename_prefix + "-" if filename_prefix else "") + + VOCAB_FILES_NAMES["vocab_file"]) + + if os.path.abspath(self.vocab_file) != os.path.abspath( + out_vocab_file) and os.path.isfile(self.vocab_file): + copyfile(self.vocab_file, out_vocab_file) + + return (out_vocab_file, ) + + def apply_chat_template(self, + conversation, + tools: Optional[list[dict]] = None, + tokenize: bool = False, + add_generation_prompt: bool = True, + thinking: bool = True, + padding=False, + truncation: bool = False, + max_length: Optional[int] = None, + return_tensors=None, + return_dict: bool = False, + **kwargs): + # Tokenizer-level rendering reorders tool result messages to match + # assistant tool_calls, normalizes per-call arguments and response + # schema, then encodes the resulting XTML structure segment-by-segment. + is_batched = is_batched_conversation(conversation) + conversations = conversation if is_batched else [conversation] + image_prompts = kwargs.pop("image_prompts", None) + if is_batched and image_prompts is not None: + raise ValueError("image_prompts is only supported for one chat.") + + # by default set thinking effort to max + kwargs.setdefault("thinking_effort", "max") + + segment_batches = [ + build_chat_segments( + messages, + tools=tools, + add_generation_prompt=add_generation_prompt, + thinking=thinking, + image_prompts=image_prompts, + **kwargs, + ) for messages in conversations + ] + + if not tokenize: + rendered = ["".join(segment.text for segment in segments) + for segments in segment_batches] + return rendered if is_batched else rendered[0] + + encoded_inputs = [ + self._encode_chat_segments(segments) for segments in segment_batches + ] + return self._format_chat_token_output( + encoded_inputs, + is_batched=is_batched, + padding=padding, + truncation=truncation, + max_length=max_length, + return_tensors=return_tensors, + return_dict=return_dict, + ) diff --git a/tokenizer_config.json b/tokenizer_config.json new file mode 100644 index 0000000000000000000000000000000000000000..418aaacbd996a6674efa89d26be06dae11433e1c --- /dev/null +++ b/tokenizer_config.json @@ -0,0 +1,157 @@ +{ + "added_tokens_decoder": { + "163584": { + "content": "[BOS]", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163585": { + "content": "[EOS]", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163586": { + "content": "<|end_of_msg|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163587": { + "content": "<|open|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": false + }, + "163588": { + "content": "<|close|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": false + }, + "163589": { + "content": "<|sep|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": false + }, + "163590": { + "content": "[start_header_id]", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163591": { + "content": "[end_header_id]", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163593": { + "content": "[EOT]", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163602": { + "content": "<|media_begin|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163603": { + "content": "<|media_content|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163604": { + "content": "<|media_end|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163605": { + "content": "<|media_pad|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163649": { + "content": "", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163838": { + "content": "[UNK]", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "163839": { + "content": "[PAD]", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + } + }, + "additional_special_tokens": [ + "<|end_of_msg|>", + "[start_header_id]", + "[end_header_id]", + "[EOT]", + "<|media_begin|>", + "<|media_content|>", + "<|media_end|>", + "<|media_pad|>", + "" + ], + "bos_token": "[BOS]", + "clean_up_tokenization_spaces": false, + "eos_token": "[EOS]", + "extra_special_tokens": {}, + "model_max_length": 1000000000000000019884624838656, + "pad_token": "[PAD]", + "tokenizer_class": "TikTokenTokenizer", + "unk_token": "[UNK]", + "auto_map": { + "AutoTokenizer": [ + "tokenization_kimi.TikTokenTokenizer", + null + ] + } +} \ No newline at end of file