From 7901470e63526a83606a3dcae2379ac62cc6538b Mon Sep 17 00:00:00 2001 From: biondizzle Date: Wed, 3 Jun 2026 10:53:41 +0000 Subject: [PATCH] doc clean up --- TEMP/dsv4thing.zip | Bin 20388 -> 0 bytes TEMP/dsv4thing/README.md | 26 - TEMP/dsv4thing/config.json | 35 - TEMP/dsv4thing/convert.py | 168 ---- TEMP/dsv4thing/generate.py | 155 ---- TEMP/dsv4thing/kernel.py | 536 ------------ TEMP/dsv4thing/model.py | 827 ------------------ TEMP/dsv4thing/requirements.txt | 5 - .../CORRECTNESS_BACKLOG.md | 0 .../DEGENERATION_TESTS.md | 0 reference/official_inference/README.md | 1 + 11 files changed, 1 insertion(+), 1752 deletions(-) delete mode 100644 TEMP/dsv4thing.zip delete mode 100644 TEMP/dsv4thing/README.md delete mode 100644 TEMP/dsv4thing/config.json delete mode 100644 TEMP/dsv4thing/convert.py delete mode 100644 TEMP/dsv4thing/generate.py delete mode 100644 TEMP/dsv4thing/kernel.py delete mode 100644 TEMP/dsv4thing/model.py delete mode 100644 TEMP/dsv4thing/requirements.txt rename CORRECTNESS_BACKLOG.md => archived_plans/CORRECTNESS_BACKLOG.md (100%) rename DEGENERATION_TESTS.md => archived_plans/DEGENERATION_TESTS.md (100%) create mode 100644 reference/official_inference/README.md diff --git a/TEMP/dsv4thing.zip b/TEMP/dsv4thing.zip deleted file mode 100644 index 2b1685dd27fba9b51d7a303181b34323c740079a..0000000000000000000000000000000000000000 GIT binary patch literal 0 HcmV?d00001 literal 20388 zcmZ^pQ?O{unyja7+qP}nwr$(Cxu$K~YudJL+jpOi6LGswbdRX48c+XA=EE0Rr63Iq zf&%cb&z5OW=f8gby+HvG0GK$tvA9@R+L_b4dO!mJf}8;Z0Q}dcssalDUiXV{)$ou0 zbt?SV90&ji0L}lHgZw|{D2WP+$cob0n*47AqyHK7FJ=Ct9%b414F-gsQ*{~2aNUa1 z@EbtYg8;O%LRtjCN-yFahB(-PRvYjrgkRrL$hqyx)fcfHj}zaiJH!#KeBm78BePh7 zm|(y{T>?8paf$O}lj1+{Dy!ZufF`OaQs_a-^4~ykyNy8*3m!l`!)HP_j4%nq0vbYtF`3= z_l+dp*NY^);`@?~DH2ObF}NtP-gziX4w`=LnJ^>A3#wd7ia?yiOr^#%&oUGWMTdac z==QYT$?I$6&5^`HQs5U;!5hX!HYae2e@{8*R8g&QqsqXH{D15@YeSPup5N2_gw9TjA=gwPruPLyHmo33= z{tOwv!0N_;;>rVomk98^4=JtJenGOL{wU$xP2yu((~dPO`xXt)s9Jydnwg>i8ri{f z09UQ=(40x*5*j8Y8#}4xdR%pGcrQun6r{l+psx{p#{dBU+Wv{p{|}1se*^{fe+0$Y z-pkO7=ey`Ir9uqhlp)yFmhkuUkqYfAz#*JWLwg-Lyp;3nMF0nDh zPDqeNl>M$l4+N$k+!fwZPAGLDG|4Y%s@H4bI|o?D#gyPX%+u%r$BT-}uoxvoo5O6S zRFRZv2Ew4&N5UVOvl|zGvbleK-rqt0W#|Ae^Ml<#hSGujcSBA8!%&p}G}O)1$%W3r z^MCU*?;juY|E&^CYx~2I`!4_Cpm3C8>3bT@4vf{a0BLMsz^tQ$H(@|QU4;Xs4dn5X z4r=LI$KM|2CLbEniH>G48%6Hh(V5^h&kCx+l-{Hl$y+!vGj+=nCp7gyu12Fj!l^wY!1Fn<;RSpX-dBEDHX>CRP|VlqRQ!+g9{E+F%&ITuEN7G1VEkC zLWuo)$vYLZ{=JKz#XtN$U;BCdjiIKosj#pZG_neizs1rDT~=vo|#Xs7Q%Svo)~k%-&gO+ZA=D z9ygosCTO_0qC{W0pt!Oh@I$Jjmc;2CeI_O`o6S&bT>s1vCE+@%b1$_JihZ zOf8e`Tz#8nFG%m9m{zWR0GepK$4q58%Lr%ElPv-$F|Z{7@{tiT`3*k24GDUm2>%7U z%KHAg*Z*4m{M_~TF?bimPLq1FPbT2YZ+h=;J%H7N$^Y(fIXZsE!7H-+(O;{irbbc7RdXK`28|6rQT8m!%_@ zLo1QjO@4w8hL}vIOAXTra6*ikXwO|_Jcfm6co$afhyPI^2)mStz~)O@MPGwF#86Ew z)xzjabGXIcX9V!ch9TU9Ry=qVc0~a6p8Q&We%8d%YngUAzRKnty@@DL`3l+Hi{$)!QKh>VjH?d$|Q1{Z^7|_n?9HLwB6*Yw3ZRt;?IdJ{lCQ z3<;|qPD3wj*5{B#7u3`tiYgyr!Oh?>N24A&wq`~i&=y40?=Z9IU~QNXnMpS+R;olo z6RHv+YqAYpRNnfhVwEe3Q9ql?W0r59xt@U=@VN*HyS-`FrU(kSbPrb)K({yWs@!yI zcB)KrJV|)fThcCUOg5@CAUMJ~%a&1ZYlhW!i$|3YxDLsdoVpPb_bl z;Jy2DY#qCs0D>#He#vDrc$X=BMrXK#qtdGHP6`+^H-fwa5!zlEuNxptEX_>^Ya5C{ zuw}u5KDPz`iSZ6xH}_TCi-9Y)b#K6nRgi2TK%sO9w@knRzw+9ft$QZ|v6 ztkw0bs_xKk0)jSlQrvZ1crwD$){-Qucw{&VX+3EV6#Z%7k|a790{~aFq0jtY>@62f zGwQkuU{zW2A&d``bgBtSGJnvQHgq|A(~K^foTQ9ie^LOTrV=Ea-(#|2yVkTi5#BdW zjF~Riki}9bYF&6YZO3jgV@t(wR)P_kso+t`gSveHhMd}kFNUt#?Yn%F5E(nvHB<7? z8)5m_gi8h}(UK9b6Pm}YC^BTsK!SHiXH3@dC_sD>%c0gW-Qmf6iR(}(-Zu1tFis@4 z-WqNYz>vqJ4n$Wi@P?(nY*vxxE=nD}_Xcf$_`%?-+ThGOFq8@QZ@Yj|`X7r^#vcNb zwyFZ!I5dfOUJN(h@@Wft0RH}dnst3c3ycFDH3)w?J^Qafykw=6Vlre#sr2<_Y|i&Z zx^0f+K`3}wm{m85nH_!;xMQSREN{R)^sa{5hIu|o+SU4$>b!MnJ->UpD8XIwY3opD zLtTHYPs|Wrs)I@;@m)H^y{tdt2elJrF0R^1H1m|b+;*LpeWwn0Q){vcWS^G+bSQd@P5 zuuUO{NE}AbfbH4^Us7q+598x4$^&8kwxsLeCQl)|`4>oTDw}L`Y)kNRcF!2*360*i zP;!&!!O8)QEJnN?twgs;_;uU1$HS^|0xjgp^DgF6mzI&L$&BiTSC7HDv1axR;!g0- zc6L8=+f$#Cw%!#*{qYdCc67Um)R%4oW>oCE0pLZxe)q?NfP=7mE73z@61g&@+UVSC{0w~l_;~O z7m$%WqEu~Y4Og`)n*Hs$UV}QBi-;A?5Yk(s0&DSjD5W;RVyp7+Tp)ATu^(vQN~d?3 z-plnGnd(dzd!?#y}!5d9}`*@DCCgw*6p{lWL+*nd>Q#L zcNE{70BmCkst5M~-o{GmkLnuv>_7M)MhhqjS$GqF^9ssDvRFQke#eF?8498nn~a;3 z>%Gp0ldKTnUXQo5Ije7t+s)exOYRAop}l}jcCQC7?6p~-is4b#pLV?F`jgAGW&0Yy zO_1$vD#77z(+6sKxjc_fp<8i<&z2z^Z^`*^IVrbSDR7j_`L25SJ3i?DMO@i7r!xqU z002eU|DCwZ{)4#wCqkQ>+L<~Tx|seS$W`zUZWaAoH5!*mo2{rnb!7tYqkCb6E$gl4 z7L_@jGp0byoCNE$PS|PjkUG|xM`DVHUt=!gRmQGE8R%VztgW+zse&} znouJL9)uS%$9pd^k2y)Kupwkenpu3ek$J<=5!FR0tRQo9szN7PLi73KA546UZGAVs zjqspVAuC@5EwUKWf$!*;r{Dcj~7gEI5s*-wK`w38@7G;DDrV0#zNqdQOt8r+lt(b{6T zh!o85)hXhF*P4H#KE;+1;ObN6CpqK_+!Ruk5rvO__cTm(A^OG1XtXSCn4J(ci?iaQ8PD5e_sJ>g8RtC!CIHg<PMD~sRk+N=qEZCO{i0q z#wsDzs>Q{DtrQ>!aM{An(@??FT3&u%XvXz8BdOA!Z)HMI2~W#WT~r14vc<9Z@|t61zyPnu@%@!R;>u^Yl(s9k9hw1*Bb^X6vP-Z z5zcM5IG@x7k*__?$4Kv)@9e6!)Yn@hBz-Pk4pVWAH%Pcqr9($8?v<04hzr$cl>7LEH5FG zEZ8$-$*u<#-0;zXJ4l(syq56gt23*zD!)C581e}PXKhP&kxSl_>xOa&Z zhxex2!PVVxW~UY}uOl+MQ(rLTNuhh$29r6~iXH&~VS5yw!);H=yngTR{&M=&besP- zm6U`F<~tlP@1cj(*I@MRPl!oXZnX-;!Dytu@9)DfwcB4HOt`BKvUKAxN8j)eslhKE za}H)~tYu9#vAA!AO5W;FS?CQfPc{DLzWm8G=PQ`*jnh9m^&vVs46Ut-_r@TX43|#1 zR)Ci+?Wve<1v>>Ea;lM1E%21WR^^lgA;Mpv)W?F05Nd4YX1(D#deOi;)=yZ@n4w93 zDuV`zCayB>u=jROZXep5wJ5mE?DBP-W&P3d)Vgo_m0*#x^mKK7?}wLDcNMR}N&x1T z02xQJ6&b_8I!RXM9Bd2F(LJ=K?0HibzF#H zOi);0+O%x`ZnGd5d`w~~$k5qW^u!AM23zM=69Nfijs%h-dL5~V@alYO=WOyPt*ch0 z2X!Todc;e-)mLY|5}WZC57aOQaYm`nP1#&6@nv610tgZJ`Kl$-hg$4FIFkJStn&W? zj>?zj_Z7tEY=C`&ZA4$=F2-@vvS$#bFFKn?>VC4-D$bvMd}N8-hsdP(=Gd}5TK+^5 z-4(suikQ7Z^W(dble9u!X^*LZweg;CWZ<&qauZf{I6ht*ltM@U{1#dsRFZGhx;;j` zlFXguMSEKoDqoHu@<2L{Wu^wT#n-<8NYDcEwjf;Kc`1>j#XG$lzVcbAQWmmCa#-+DSUXnuD=#H?-T z)o5F!p7|O{2q7DoYU?w2$Xll@(;|oKu*TzUfK=j*hg`XnKxQQMXWP`1sPepPv%$gg7(Tym#e z&&c@hFEUbYALY#cI1^}mgug8*6cT9NH#V}%<>FC%=vjDW;jnt|y5#+e@lu!`j6%W# zE-KC=mqxB$uhTD~4Btcxqt69MZv4iU%(z*EF{qlsxLZQgf8Cthz5Yx)kyFYui}`xW zqDAL;@y`l+e>sTpX13YmcW*>}2NT~|{BoOor_Zzy zuZz0!V|POYl6Z*+lP_rQ3kn1F#0w600gTk`lvuVDehE4gGigj9ivOwsQ=_G3S!>&} z&mK3i6M2Y$9K^#y>+t2MIQ#vlLel`2DxKaGLPK1B{hP6P(qmAlfw7T z@GosQxKww{Rx6n$5BWP4WFa__*A87$0@o3~B;HUWdaD*KCtFbm_6WS4M4p933qS5J zQ#Z!XR;zZ+&?M7pe39f`LL%V%Em~kajytTp{Q97D=f#en%delKv(Ll*6e#5C&Ghs7 z`SAT{^e@-vCt`;z`qH|?)Ae%0=RJPjj2~s#w|teWpD)B?8@b`taX8)eEPr!B4gKZ^ z-SzY9#_nTU!J*zZEn_Z1 zWP0P5^A0{(A!|$ghS28C>qU=T8Yd|Y?w5dYHcu!XCr39ov7Z;+yqKUrQ#}+CXa(VE zw@HU(92n*#(1+$g#%_)bfsl&?6F~fzK{{mkVrDl;%^laJD7lL)l9qNY=sKL?G`1`8 zbjT~yNWgwyiBEiMM|CFlo_;f4{aHg1>hP zfWe$c5jI=`;I0!q5~FVsh~W&O6Jhbac+OHp;bwT(*in*0a_M_pC%!oc&bej<#UR0O zl_@h<^HKTDdK?fsq$SuksFOc z6S-?%yv28Woe+a}?^p?234QTx)NZaXn>V^K0f$}D3L&l2E8CkZtn3SVsJ1N_=O9FFtZ;bF+F?b+P=SK1K3k`1Z zmM)@^yRD`Uhq+#;y1`w)&I_}1GAr4YE!nk24>rQ#(cJ4PQMA>5lp_@2!1K!u_JInf z@Aynts>SJ5c5AW-KUz%jI)o7(8o|-1MdO!ba66T&p`^-2;k|2eX{tyMGA~;6BIF3R zsHZ8FHhuZ)?zO()_dsXC=5rdn)a z#9-Zy;nlW%Nl;*j5C|f{bBxG9;a_t_?bjc9U}F01HQEg{W6Q|-2=S&x<#_u~YnwfU zg8Rp)OO|*4{cl+<`-A~A?fkrA>k8qa0pD6(yeXv;&5HDZ1p*ci)&=xp&oN!QdDsYM zi`v*i^J1NRDE{{ri+o_7OwM42cuer(U>=63h;t_eFE!?Zf%7&0=$aIs z@$c524hnf=|1fJg;a!W(DRw>2Kb439DC#;MA*_JPYBGWw(`-D6cF*Jk$}~d^F(PXe zDO3Zf&A~3nS-QY$M9cw(GR9~h?H(_`)rSu=$kdS^(1=QH42S@LCEZFEVm%MPyC4z_ zaGc=VtmP z+1tEg9G^_kHiCG&VlVv!))soM_CLj*{tMmwFw>Xc<2Nu@Hvnpu&u(s9pFO~vGVQI! z_TOe7FTWYa0|TF_YpUWIbMs6Ri>L9zb!=^u_*wePPCkzQfY+PO)osp&R07_EHT^F> zC$TccQk127LaJi(%m~Q`PV>4iJ43QEJXJzURtjBzWvP~ogq0eS(0XUr!h7m>tX6KjM7E>fC<|G05Z$!HT)FOQO#-uEz!NA7_%_d}LcVwo)~_5Z zHCG2;ZWy69`{@L(rn6D`l>hl5w~8;Q7(EYD$nA4==n)g~f9EPxKZ;S|oL=*#zQ;xY z!^9GneN&zHd8X3Gpncxr(qd%_J#@`an^u#dZWWycA+LkQuf_H2yuTVW)A1XO`Rb&| z7MD^DcHPP^4MM6Fd25A^9GlJ!PDk%=i3}$k zwO@F>k`m`H*0ajVQt&$_T!p}BwAVroBcIv27e-c}_UEC!9|%)}v^J$;<^GVR z#1D=zcfkhn#NpyO%~ZcyTOwaLVZa_1jrl%AN+qkNEHtA3>lR$J$zrVnd;O@)8W<@@ zBYP?_+(cO4@V9r^_kqtY!3W z@|gEM^3N9Re)_JRCs`=Ll>pa~v>7`N2elehka3LBnWId%g^V1D1bV{6B^XSn_B;&z9^pH%qGCdV^c^#M3x5+Ma38_GEx4vrD)OTPDalj*Q zPODxGIb#=riY;F%J#H`7T{gzScdOw9UiTN*rS(pLFnuu++;K{^`3lRFVc&d_G;9Qo zM-CWd0f%!a9hT;}GxH11+61pk74l>A;0p$hmSd%xCmu?*EGGGH(6GE;RZu#`@p3;b zHrOWCiJ&=@L0kX`0F+0}zNi2~?0FK8$9qcY8QkqaEMrj;P43Z~)Gil%kByfO%!SE2 zRWVy%DmVCOK>$_+>w}X&6xN(Lx}3fSeR>2Z-TBnmu|JXpoPq!tL*guzeyLfBoD5Ag z^HV5G0BR^T(7O+cQW9Mugeypk8#7Pi+-&d|Ji2j;;N4KPhP+t`O@x^Viieb}MQCdb zh;63&#%UpMWwUr8E*fhU&)?Ekl7O#5(y_HzHP5BT6TpTy-!Y7ZZMgMIzfpPF8nAfr zsQ%rz@2=n^%6(pG{WwF}HSx~m_&Z2er1ud$;ff2tlKM*V)U1Hy;AzA3(fU3KNK`6j z@5-l`nC%1E<8fFOa2eddn$);(wITnVWpmBnVE{7n34VJhAKgF1(utHYZC84viE6a?jF|m})i(-iI zVcEfv5kky-W`X#x+6FeePQM7~_G`;cN6{%e_A!~6DPe+;b0*Er-LbNa_bCp4W918fQQ0fCKJPP2{C1 z7Q=*ZSi%a$eMdnT0`F@znobRO6>L$f?~1jmzd(|)(9 z71(3UA0Y}&Q^lJaP3JrY;jI@bSyr<5suv2;^&Y_1R&IBM6_4{zh!1NU&F=qC7Bv%r?4x)Re%@0*s5HJz+h z>y$&ihc}bHC__Avz&IDpvC2r7U!mjlc?{~hM2TgY_IZd}vm!2eG&+@iLXDpHQ2u

@M%^@q@SWQ-21DUt3ru&)jG?p|**i<@*o$J@Mt?RpC_bK!PoY?sG@Rg|3n5lTQL==BxnD6eRBkJ6mSpVlIZ zsShBuqux}nsw|h=sGd12+dEIhp>NxK8Ki~0RbnQI=CUUgQnh~c7 zUEcR4)<7iad)Y!-6BR!vm8fWOR#37YxlAiGsJ92P%_w+LISOl$Z+14QGT5bU42yap z9G*CNz>}#5=4&~_XiD$HenRnUYBe{F7}~QYsKO|#+oAkca5-9>sN@2aeb?b3ZPcIi zO-Ngvi0`K>NpM2Y8d#bMD%#LQ68%pnPXi`}&tZ2Q5jbo||8h9zd)nxvOc6Ad+#d?NaAC(Ragt(;nMV_i`FH6ba%FSgZwyV&ak2zC$Y#=cE{dOtq8PC?&n+`#-{4~jgURj|PjhX+&BMs9riv z*luqWC7{kn#B4et<@>vR7==$R*>1EsDQj)G4Q~gNm8gFcjjNs;T=Q<60jXZP)an9V zXULx|#omClq0u{jx{zlX`UhLhEhsFb=)^-A=2H=UQpzF7lgQT!N_rRzRL+u+I%a5> zkSf>&^oA5_)eUl5Nh=Xuh81~JD#0&wVd00JiguZ{Uf1)yo~Wu7+CRZbY3ur0O6pS; z@%I(A;+Idfq>|-^SMSj}Qb=PEls-=HQK_iDcH4pdqpMa!_YvK5x80}(1_HSO;Pv(^ zLNkh#MUD6+z!NMkHk^{)D75=E z?0LZhxCKPvQpx>Jt8N{wFS!9UqJX|+1zeq4kVYtM24XZS&on}bfoZLRv%<4+hsfgx zmi>p+t$l4y&C)$uOH-*T?&n4l6pErIGRgc8+6MdbPi11C=`np}KU23QroP%+SI($! z&B9hCtlJasceh|gZRG6_w3X}UiPrYdZK0a_9CqwBK?h<)7U}|eY8Mi+AG}GC2f)Bg zMnqI~mo_pINgdRw0qAEH&Un!T58XAtb+A8q%K+KI4}Cf8kN&?ptevMf9%Cc`fNr(_ z-eFn)M~8*Drv+Q z$za3_V%^S@=AMfy@5i0fW*Sla@EBLFYU!EtP9hzmP*wOD7rUzd>sXC^Ur092T zW~nQ(X~nkMBSws}S*nR`lF8|-(Uo-YaBw%3of!kc`b><7hrwL5b^DX!1Ez|ybfjKIo^J=SlPN=MF#2f>zNB1DzvF1qZ&i2Y@LYOZco;9jA^4dghmJ_)1dv6`H zO-5pUC!7gFhN#bhe;r`}1jBaD_3wpQC!^S(Rg30@j&<{eSNcy$KpW#b&)T?UJOc#jJhF9+1-SO6aT`vx~>1bR8Zx=__Ky zyi3_O-uI7vzK#uP-~yT8>GO$v3$n6UWNeEO%DZOtXyx14{7XL8!OTeJ=PiMXC1E2* zU|jrJ9+#sR?=0-9EWOUBTzx(^7uCsmDty*6&l)F?JW`e=%bR7Qp78E!p?+lfbc?6| zQgbDWPP!eZ#;6w285B2zDM&Yi`gibZS=`06&h|FA6*=M~+m6 zi2~J{8a7m%)C;pmjFI(eslCVa?g&^@Xrf+*C0_CnI3Vv{!X~Rcern z_D*RldoU?2_DG&%C*azAfOMqfPL5L&gCd&JkJ^uvHEyb=y@etrnyS(#+bnAwj)!$e zs5*fH>v>|-w8HGs(VH*lh>LC6Xwe$Nv%|80hygQ5a^wh;4yn2zLBHLa#yscrkpo8t ze4@0sKtq;+cD2h+@9F3En-8&o(%eG5NB3#U!zgvH7+0~T&Ia;TZh*;~rq4l5WT^%E zI9lbhi@KgBts{PAnPPv)?ItlOI%oDNKn3e;s8w0}`^F2@G6KpZlitas>lREUFVG1h zv61){PIO)R6TDx__X|Z6QSZ>vH?6~nNPsBeXk2yHlO)oiw6B{{^hQJZ5r&hF1mTcS zV2nx9=qBe~j@Ep858$up-c3NkA`0KO{TJeR5sbdxr<^dsx>(cs-Us$!v zWLuv-Fl7(wh2}l|RcDOyF~_h1 z<6sCbR|o2f1Si^YF^B?~n8;ci0lH8F>QIFGKO3vu2n;u9`5PXZFb&u;@qb}Xu8L~y zLVNY{4DzX-VPLS6AURI!0=vZ=iT{egCxwd%_Xh8 zaX_#yu~3^g8Y1j8@d9j{zApnwVpu$plYSDbQ3_~U-4-FlXw(e@u)v5S3COq$OePVr znJY2BLzw6+12g(u^DwvoLZ!Co)eWzq?1M+5HgqEDgPw9C{Dy|skpnKl;7CMN#A1LI zw4=^8EJ9De-8sYe?ZiIWMX)u*4h-}*b-)LZfdpOfF>y0?vLjqp>#mlxgTn<|&>6Uw zx*qMbqwyg8o}1=7!qeGotq@a2Lu3aL)Mlzs2Ua7jX`+Rq5<#>;UED@! zfdQBdzaaKWHKtFZKRz(Sye=1oZJL<#QM5cjNL>Swg7lpx1sokB33^dSfi%l2!ufsM zeVN8c;|F`P2Gi4{5(4T2J{AG0R_HPXZ+M=pF@p!XSm$>YWd-_y#vW_9v2kl+A3i*VVDv6k zfc8Ga@e&u|_4o4v^bgtPK@)XNm7hOZa_u;{C;@*5PKt$o3;b5>_D5(kc3f3 zMiB%7B}ohX3p_a$HJ7Ke(o=i|-SY*W31pa`lkiV@FS5klhk1*={23sv> z^6P_8@1Hc>$oXKd4WZpg`|%oKE+U8zAYZ3-;KIOZP%^A*6RAG~GKVmF@w*r2CA_i-q10?a~; z!M|ZDY?+>_$!2^UwQ0jVoU3&rnKQZWBM5ydo*P&DRE)NQCU%9T`Q5lt@XTN$79%qJ z3-34z9Yrn5gSP(UrCI=v zz!2a_XfUi076?`EgT^^Wva%>7LDR_Pd})LMNC_}V-lmxuT*|$%iwq25PlVwm!0zGL z*#|vBAZxnX}C@&XSH{)tF60 zeHcg0O$H9ifbk$dIzotxjqlP&enh^lFTC?av}xc0iB>A&w=uV|q8fc8ddO&+{LH`5 z*S)$z1^Bf|M9mExf(!jcuUiv^!!@8;)8ynuPhrwr8I@8P3&Myx-B`H+3YdI@o}Ym_ zb3_NM;N@f@efd5xXPp57NVMv3s5Sf8{fD#^u5+o#sYk}a=q?P>rJn+t?j0)qAkGLWF@o1Tz{G0nvYxQ)6g6gd3=^ z?t&r&5Vjrs524wW=0LWYCn89jvq#P}U?>Mf%$2}&HwnU$4@c|V-;e|X!Z8XNkW@%E zRXh$lbqb+G=+NFS*Wj%$6Bj2YHV>mYjAh})$To=ykK1Eu0|`O$=06ZA{+S)~444FM zzytLm2y~8U~WCry%iPSfverL4SSeGI6g;APN}I0PhUx_|*`ba4C zg|FA#BwHNL#X}_^{4l4M0YL8~&)|gs7&U1r@f>Wrq|phMAz5jTrvQBuiw?*6F21Q` z83r7SfvPdaA;#$2_~iA?_!;>aRbn#4n-G4=X<;ZZljo!$cXpDpk#&}saOKyT8L0tu z9Nn2bVapF4t*#bK=FEwM@bJ(2Y&c>v77azxuOWrzhugu)#`W0ZjfuR4&SvPl(;r6z4aQ#~q- z+QvIz#Ok1Xh|SdEtShadqg7pA{v>g-yt}a5MwNl2Tq^MX=iFS-TJC-q$R!U+gF$K? zQ{otkmO^R^#os--=qK)=eZ17#;5}p53D#ztttDSFv35)bRfZvXTzbvP`bs%X35x-l zC)>Fi$FxN*Pl(N5vu09L>$La$=$%Gss*a^1ts#A?WHI(Y0ZfZVG^4DSXK6Li7{e6@ zuadt!zqb?h&k{T6i09b~eb2|orGvRLI_2Do?;fbEXC_yRLR9ZlS-{y~2-Zg-beX_J zu_nK>0A72|;_?bb+@|-GK)mkJ3#XL|zxct;N1Dwew-TX0t^8tQ6TT792qZ>yGeDF1 z%!l+0cK;5>3Z4KSaKHH$7+Hp z?ZnfQmfmiofm4gr`Mnqk@Z@A&f-~wleTKCh);6|^p`z;2b&ct}0El--?yO)N`JWkp zQ3IK3G>YCwCoN?U92q2Q(m$c1qX{|<%qV8IZhmOd%!Nua&*zODYbwUIAiSkSe5)Xs` zfpIhBf+y_P%p9F}dSr&Lexo2O$e&h60(2B$ZLwp=_+d{R%NOpl-k8!1SHV?@OaCw3 z>mHNE4=-?tjXtvO@}|D7VjnYy!x4oLMiQA}6S0SYk^p4L*u+?IPsL@RH+qrDkGcK<)l{7sF@Q zt?`QMysrk3S^Bf{vtr5-aLp{+nGTt#_g-<+?E`O@glaF~Zl&{;fc6$h*rI z0^VNBHB*Tjmy{jXXXWk#$KSP0{29^5{C$l%*yzLWqzN2tdtu=&?NPQrpp1kjr;zD= z5g{&26lYwXZLq2BeGnC26Z_Fnl)Y`?4|8>KR*J_#%I;?+mA?P3(kvR?zFs1-1@at% z%^lSUwjL5bn&s1x2AM=5U^li56eXtjfD4cYJi1J6=T?D=heR1KL|a`!>X|x{JKR|~ z4e(e8#B;iK7)FIXD06G?A_PZkj}YhY_n7#Xmr>3psRi#JZaGyD z{%h<5lBPx@7z0d3B1RMH%mK#@?P|>F)%hO+cu~7FiTMsB{E&wp`ONS_gPQb41eX?H zJx6e<9x%Yjn~Z@I2aX8sGXieS1-w2zj0QKBdj1*bs3%97k!E1sS_1iUX|b+)8Iqgq zk6-%ZidrLZsh~@ty&vv6eZ54cW%C@5%OAEW4QW-m z3_qWhTA*I<;<9M@;v4mXSLHmMqCvIIZdrYGwVhucEfm-*UL{IthsPvrwr{X~$#uI` zyWX-aewq;pObUy$SV|>DquN zV(~_EVtWp?5YLq6<>FP6V5KcQb`5{|n$Ez#Ed_u15HxLq(WTZ$tcSNt*Y=(UJ7a`u zC|(Kk(*q)6jKZWetaOQ48Lp47Wcv!{gOKcjk7SOw(9BL_(WhH^)VQ>B7>)D{V)+1N zc4k47MFnThDu2qLO1eHx&3cZvO1rqNudj{erpN=#a$v2S4xAqq~rgIBZM29U5ckrnqFnX6fu~#J)iYskRKSIqQ+S zdWa4N*QB*78^|LCZ!+qu{QXQAA3k4V=ehR^RaDtU$TuS=F^;2%0eSUtNFlw~7Lxc6 ze{Tdb`F-X zKgq1tb0ZpWH?B`dZZnb@JVK?x8+s{usaK6pvx84eo4 z(H5Xzr9~+2L``xAO3-_>+}Vl8VR^fcs?zjPRqj5h1cSwCXOEli$rj;;ER@2h<-6AF zs-OSa<2Q108bj6T86LD`m}4%e1K-!4KRUn=8``J)(=^S=)UnC zQ?HO)kAryfKNR=A8bOfafQbI3cq)0Cb za1cUoA_7tZN;fnqp>vZv_sj`*uIGL;d*=J~%-Va_`uDDfE@$|8c_=5Ddn|1516wvZ zd%_`=xpg1=j>J}~sw_Q!%P|qETs`bp^L&TV$<5;fgOjeQ=}+F}%5&d+hwc{$XB$22 zr$Q;;DoLes+&3TfrxpFV zw(dnQhN{{mbh~L*GvGfpJyfK9R~qA48Z`s1RlkY~x{gYk%X{e6ZkZLJmdUiEv`1@v zh>YN{YBa1d0Ksl=s;vR3v#!vkOb-$PJ8gCCLzH+J~M!VO%{^Tcny-_IfFR0Su``_pNE1=#EEkXt4m1tH+X(3T9xLWMD4 zDfEng{`q0qg#Gl0y;+cwBKM5hPTFOE%cAep!qXG}=^9K+)$%WT5Ze3s`2>7|nq^h9 z8pVg9!7TC!Fb)6a*X%9-zKNc-cT@YhCap_l$APmoBaJB)F}wQG13@>^lPJBTvMX&0 zdY>9*6nT&Ajon*Vdks{GYyO1V&J+5oTMV;>14?N=WU2UQHg1Zm+=`QLTsENfpLZm6 zxn=`J`|^ssGKzU1JlA!=FWuN3AErdDLc5b1m~e-ug|+;iKK zJKwdH1)$`#dzio3TiltYN$V%*Pb3cSypvP5Em?*N;vV)LYd?Q|&>T?8HcL_4{_`Z4 zde*bL+fH0p>!6#TAR-1PD@Xt8HAxVRuM&eCe8P{slkI0Je5nQG@lIc;X=X7jz;e0` zP`(BB2@o<2yT4&!Hs{=J{D{b(&$10^_zG_KGe_C~(IfG-W;Wuez9tVFVxXbZ^>%%DFp5ISM^hb83m7*KEp^jU` zR`$}ru61xE)KM)Y1Kb+y#vLszoK}M>kbV{7{%fObw6-h|0n5~HzHxe+2)Z+}`Oz80 zO4_#}*Yk4vYE87_7xl=`PC6ql8DMXRxy@{3^Ab3f&6w7JA@O5;v&y`k63sPTlL9pw zbzDSLPd=IcT;MSKQBkk!wc4=P7$ed+(L|5xJPDI|?1GMNW07#6>ZCDx>EEE{*S5DXr&w5&eob`;x_eM^EqOVQdOU) z@po|u(lHl&<_k>@SMCb+>t~nG_~vPG^91;=EzV-5glWGeL2AP{*o!>?w!Iu2SW#gX zIi^u2;46EDQDKb{em~#%@!1c%#sG26OwcSljO&V3_*^D~QdbwK{fM9q{Jw-rL4J?4tkm7FkSFwU@3i+;x2sBwcC;jxf9Rl-@`CH5Uw1Xr*RgJ; zH+8B+^2m2R!?cZk&!eug5DK1p*b;lc)(>QKEX$yfEUB}4IyKAPa7;$@TwgUzSgdhrcxD?j>2GQ$J8)-gAS9B7o`SxkNt&jEs|lM&Gfe^=oQjZu?y8^jL&BGm zsNR;rH7bWuLE^SP!e@}7vvswr!+w7nZ9L`+#73xL`Ax)Tk*>To3?atdJLUV1t^Srf zaqq*eJEx{QPpQn~ek2l!Wc9=tR_`lEJ>DISwvN#%m&V>6Xcid; z_V!TJa_<)<7HW1$p1bW3Kr6HY=1 z2{G-Mg>U3Ea=A=#T>0f-*3qdYP|mS*C8Zmsu^GzT3VOf^yq5RzmI4f!RA4Xs#5p;< zLn+DXL{8mSt)hlIG0s_(Ibvgg?8|`5j>AmB3n{gq8w88*wd@2t9{F=YMSnlc>|^CB$`)@oY0FQcHK32qmo2W;1put1;%89)+XlvZA%o%$e(vSq4t55+czXQd z3Gw`2MxSzQ4h45Aut&a6lcMc96<`DrMQ$pvu2#X7rv^L-;lY$Nt?r8!Bvjk6Enqj8IS_;IL<4~@KvkN*7{9~O*T tuple[torch.Tensor, torch.Tensor]: - """ - Casts a tensor from e2m1fn to e4m3fn losslessly. - """ - assert x.dtype == torch.int8 - assert x.ndim == 2 - out_dim, in_dim = x.size() - in_dim *= 2 - fp8_block_size = 128 - fp4_block_size = 32 - assert in_dim % fp8_block_size == 0 and out_dim % fp8_block_size == 0 - assert scale.size(0) == out_dim and scale.size(1) == in_dim // fp4_block_size - - x = x.view(torch.uint8) - low = x & 0x0F - high = (x >> 4) & 0x0F - x = torch.stack([FP4_TABLE[low.long()], FP4_TABLE[high.long()]], dim=-1).flatten(2) - - # max_fp4 (6.0) * MAX_OFFSET must fit in e4m3fn (max 448) - # 6.0 * 2^6 = 384 < 448; 6.0 * 2^7 = 768 > 448; so MAX_OFFSET_BITS = 6 - MAX_OFFSET_BITS = 6 - - bOut = out_dim // fp8_block_size - bIn = in_dim // fp8_block_size - # bOut, bIn, 128, 128 - x = x.view(bOut, fp8_block_size, bIn, fp8_block_size).transpose(1, 2) - # bOut, bIn, 128*4 - scale = scale.float().view(bOut, fp8_block_size, bIn, -1).transpose(1, 2).flatten(2) - ## bOut, bIn, 1 - scale_max_offset_bits = scale.amax(dim=-1, keepdim=True) / (2**MAX_OFFSET_BITS) - # bOut, bIn, 128*4 - offset = scale / scale_max_offset_bits - # bOut, bIn, 128, 128 - offset = offset.unflatten(-1, (fp8_block_size, -1)).repeat_interleave(fp4_block_size, dim=-1) - x = (x * offset).transpose(1, 2).reshape(out_dim, in_dim) - return x.to(torch.float8_e4m3fn), scale_max_offset_bits.squeeze(-1).to(torch.float8_e8m0fnu) - - -mapping = { - "embed_tokens": ("embed", 0), - "input_layernorm": ("attn_norm", None), - "post_attention_layernorm": ("ffn_norm", None), - "q_proj": ("wq", 0), - "q_a_proj": ("wq_a", None), - "q_a_layernorm": ("q_norm", None), - "q_b_proj": ("wq_b", 0), - "kv_a_proj_with_mqa": ("wkv_a", None), - "kv_a_layernorm": ("kv_norm", None), - "kv_b_proj": ("wkv_b", 0), - "o_proj": ("wo", 1), - "gate_proj": ("w1", 0), - "down_proj": ("w2", 1), - "up_proj": ("w3", 0), - "lm_head": ("head", 0), - - "embed": ("embed", 0), - "wq_b": ("wq_b", 0), - "wo_a": ("wo_a", 0), - "wo_b": ("wo_b", 1), - "head": ("head", 0), - "attn_sink": ("attn_sink", 0), - "weights_proj": ("weights_proj", 0), -} - - -def main(hf_ckpt_path, save_path, n_experts, mp, expert_dtype): - """ - Converts and saves model checkpoint files into a specified format. - - Args: - hf_ckpt_path (str): Path to the directory containing the input checkpoint files. - save_path (str): Path to the directory where the converted checkpoint files will be saved. - n_experts (int): Total number of experts in the model. - mp (int): Model parallelism factor. - - Returns: - None - """ - torch.set_num_threads(8) - n_local_experts = n_experts // mp - state_dicts = [{} for _ in range(mp)] - - for file_path in tqdm(glob(os.path.join(hf_ckpt_path, "*.safetensors"))): - with safe_open(file_path, framework="pt", device="cpu") as f: - for name in f.keys(): - param: torch.Tensor = f.get_tensor(name) - if name.startswith("model."): - name = name[len("model."):] - if name.startswith("mtp.") and ("emb" in name or name.endswith("head.weight")): - continue - name = name.replace("self_attn", "attn") - name = name.replace("mlp", "ffn") - name = name.replace("weight_scale_inv", "scale") - name = name.replace("e_score_correction_bias", "bias") - if any(x in name for x in ["hc", "attn_sink", "tie2eid", "ape"]): # without .weight - key = name.split(".")[-1] - else: - key = name.split(".")[-2] - if key in mapping: - new_key, dim = mapping[key] - else: - new_key, dim = key, None - name = name.replace(key, new_key) - for i in range(mp): - new_param = param - if "experts" in name and "shared_experts" not in name: - idx = int(name.split(".")[-3]) - if idx < i * n_local_experts or idx >= (i + 1) * n_local_experts: - continue - elif dim is not None: - assert param.size(dim) % mp == 0, f"Dimension {dim} must be divisible by {mp}" - shard_size = param.size(dim) // mp - new_param = param.narrow(dim, i * shard_size, shard_size).contiguous() - state_dicts[i][name] = new_param - - os.makedirs(save_path, exist_ok=True) - - for i in trange(mp): - names = list(state_dicts[i].keys()) - for name in names: - if name.endswith("wo_a.weight"): - weight = state_dicts[i][name] - scale = state_dicts[i].pop(name.replace("weight", "scale")) - weight = weight.unflatten(0, (-1, 128)).unflatten(-1, (-1, 128)).float() * scale[:, None, :, None].float() - state_dicts[i][name] = weight.flatten(2, 3).flatten(0, 1).bfloat16() - elif "experts" in name and state_dicts[i][name].dtype == torch.int8: - if expert_dtype == "fp8": - scale_name = name.replace("weight", "scale") - weight = state_dicts[i].pop(name) - scale = state_dicts[i].pop(scale_name) - state_dicts[i][name], state_dicts[i][scale_name] = cast_e2m1fn_to_e4m3fn(weight, scale) - else: - state_dicts[i][name] = state_dicts[i][name].view(torch.float4_e2m1fn_x2) - save_file(state_dicts[i], os.path.join(save_path, f"model{i}-mp{mp}.safetensors")) - - for file in ["tokenizer.json", "tokenizer_config.json"]: - old_file_path = os.path.join(hf_ckpt_path, file) - new_file_path = os.path.join(save_path, file) - if os.path.exists(old_file_path): - shutil.copyfile(old_file_path, new_file_path) - - -if __name__ == "__main__": - parser = ArgumentParser() - parser.add_argument("--hf-ckpt-path", type=str, required=True) - parser.add_argument("--save-path", type=str, required=True) - parser.add_argument("--n-experts", type=int, required=True) - parser.add_argument("--model-parallel", type=int, required=True) - parser.add_argument("--expert-dtype", type=str, choices=["fp8", "fp4"], required=False, default=None) - args = parser.parse_args() - assert args.n_experts % args.model_parallel == 0, "Number of experts must be divisible by model parallelism" - main(args.hf_ckpt_path, args.save_path, args.n_experts, args.model_parallel, args.expert_dtype) diff --git a/TEMP/dsv4thing/generate.py b/TEMP/dsv4thing/generate.py deleted file mode 100644 index c35c8030..00000000 --- a/TEMP/dsv4thing/generate.py +++ /dev/null @@ -1,155 +0,0 @@ -import os -import json -import sys -from argparse import ArgumentParser -from typing import List - -import torch -import torch.distributed as dist -from transformers import AutoTokenizer -from safetensors.torch import load_model - -from model import Transformer, ModelArgs -current_dir = os.path.dirname(os.path.abspath(__file__)) -encoding_dir = os.path.join(current_dir, '../encoding') -sys.path.insert(0, os.path.abspath(encoding_dir)) -from encoding_dsv4 import encode_messages, parse_message_from_completion_text - - -def sample(logits, temperature: float = 1.0): - """Gumbel-max trick: equivalent to multinomial sampling but faster on GPU, - since it avoids the GPU-to-CPU sync in torch.multinomial.""" - logits = logits / max(temperature, 1e-5) - probs = torch.softmax(logits, dim=-1, dtype=torch.float32) - return probs.div_(torch.empty_like(probs).exponential_(1)).argmax(dim=-1) - - -@torch.inference_mode() -def generate( - model: Transformer, - prompt_tokens: List[List[int]], - max_new_tokens: int, - eos_id: int, - temperature: float = 1.0 -) -> List[List[int]]: - """Batch generation with left-padded prompts. - - The first forward pass processes [min_prompt_len:] tokens (prefill phase). - Subsequent passes generate one token at a time (decode phase). For positions - still within a prompt, the ground-truth token overrides the model's prediction. - """ - prompt_lens = [len(t) for t in prompt_tokens] - assert max(prompt_lens) <= model.max_seq_len, f"Prompt length exceeds model maximum sequence length (max_seq_len={model.max_seq_len})" - total_len = min(model.max_seq_len, max_new_tokens + max(prompt_lens)) - tokens = torch.full((len(prompt_tokens), total_len), -1, dtype=torch.long) - for i, t in enumerate(prompt_tokens): - tokens[i, :len(t)] = torch.tensor(t, dtype=torch.long) - prev_pos = 0 - finished = torch.tensor([False] * len(prompt_tokens)) - prompt_mask = tokens != -1 - for cur_pos in range(min(prompt_lens), total_len): - logits = model.forward(tokens[:, prev_pos:cur_pos], prev_pos) - if temperature > 0: - next_token = sample(logits, temperature) - else: - next_token = logits.argmax(dim=-1) - next_token = torch.where(prompt_mask[:, cur_pos], tokens[:, cur_pos], next_token) - tokens[:, cur_pos] = next_token - finished |= torch.logical_and(~prompt_mask[:, cur_pos], next_token == eos_id) - prev_pos = cur_pos - if finished.all(): - break - completion_tokens = [] - for i, toks in enumerate(tokens.tolist()): - toks = toks[prompt_lens[i]:prompt_lens[i]+max_new_tokens] - if eos_id in toks: - toks = toks[:toks.index(eos_id)] - toks.append(eos_id) - completion_tokens.append(toks) - return completion_tokens - - -def main( - ckpt_path: str, - config: str, - input_file: str = "", - interactive: bool = True, - max_new_tokens: int = 100, - temperature: float = 1.0, -) -> None: - world_size = int(os.getenv("WORLD_SIZE", "1")) - rank = int(os.getenv("RANK", "0")) - local_rank = int(os.getenv("LOCAL_RANK", "0")) - if world_size > 1: - dist.init_process_group("nccl") - global print - if rank != 0: - print = lambda *_, **__: None - torch.cuda.set_device(local_rank) - torch.cuda.memory._set_allocator_settings("expandable_segments:True") - torch.set_default_dtype(torch.bfloat16) - torch.set_num_threads(8) - torch.manual_seed(33377335) - with open(config) as f: - args = ModelArgs(**json.load(f)) - if interactive: - args.max_batch_size = 1 - print(args) - with torch.device("cuda"): - model = Transformer(args) - tokenizer = AutoTokenizer.from_pretrained(ckpt_path) - print("load model") - load_model(model, os.path.join(ckpt_path, f"model{rank}-mp{world_size}.safetensors"), strict=False) - torch.set_default_device("cuda") - print("I'm DeepSeek 👋") - - if interactive: - messages = [] - while True: - if world_size == 1: - prompt = input(">>> ") - elif rank == 0: - prompt = input(">>> ") - objects = [prompt] - dist.broadcast_object_list(objects, 0) - else: - objects = [None] - dist.broadcast_object_list(objects, 0) - prompt = objects[0] - if prompt == "/exit": - break - elif prompt == "/clear": - messages.clear() - continue - messages.append({"role": "user", "content": prompt}) - prompt_tokens = tokenizer.encode(encode_messages(messages, thinking_mode="chat")) - completion_tokens = generate(model, [prompt_tokens], max_new_tokens, tokenizer.eos_token_id, temperature) - completion = tokenizer.decode(completion_tokens[0]) - print(completion) - messages.append(parse_message_from_completion_text(completion, thinking_mode="chat")) - else: - with open(input_file) as f: - prompts = f.read().split("\n\n") - prompt_tokens = [tokenizer.encode(encode_messages([{"role": "user", "content": prompt}], thinking_mode="chat")) for prompt in prompts] - completion_tokens = generate(model, prompt_tokens, max_new_tokens, tokenizer.eos_token_id, temperature) - completions = tokenizer.batch_decode(completion_tokens) - for prompt, completion in zip(prompts, completions): - print("Prompt:", prompt) - print("Completion:", completion) - print() - - if world_size > 1: - dist.destroy_process_group() - - -if __name__ == "__main__": - parser = ArgumentParser() - parser.add_argument("--ckpt-path", type=str, required=True) - parser.add_argument("--config", type=str, required=True) - parser.add_argument("--input-file", type=str, default="") - parser.add_argument("--interactive", action="store_true") - parser.add_argument("--max-new-tokens", type=int, default=300) - parser.add_argument("--temperature", type=float, default=0.6) - args = parser.parse_args() - assert args.input_file or args.interactive, "Either input-file or interactive mode must be specified" - main(args.ckpt_path, args.config, args.input_file, args.interactive, args.max_new_tokens, args.temperature) diff --git a/TEMP/dsv4thing/kernel.py b/TEMP/dsv4thing/kernel.py deleted file mode 100644 index ea7976fa..00000000 --- a/TEMP/dsv4thing/kernel.py +++ /dev/null @@ -1,536 +0,0 @@ -import torch -import tilelang -import tilelang.language as T -from typing import Tuple, Optional - - -tilelang.set_log_level("WARNING") - -pass_configs = { - tilelang.PassConfigKey.TL_DISABLE_WARP_SPECIALIZED: True, - tilelang.PassConfigKey.TL_DISABLE_TMA_LOWER: True, -} - -FP8 = "float8_e4m3" -FP4 = "float4_e2m1fn" -FE8M0 = "float8_e8m0fnu" -BF16 = "bfloat16" -FP32 = "float32" -INT32 = "int32" - - -def fast_log2_ceil(x): - """Compute ceil(log2(x)) via IEEE 754 bit manipulation. Avoids slow log/ceil intrinsics.""" - bits_x = T.reinterpret("uint32", x) - exp_x = (bits_x >> 23) & 0xFF - man_bits = bits_x & ((1 << 23) - 1) - return T.Cast("int32", exp_x - 127 + T.if_then_else(man_bits != 0, 1, 0)) - - -def fast_pow2(x): - """Compute 2^x for integer x via IEEE 754 bit manipulation.""" - bits_x = (x + 127) << 23 - return T.reinterpret("float32", bits_x) - - -def fast_round_scale(amax, fp8_max_inv): - return fast_pow2(fast_log2_ceil(amax * fp8_max_inv)) - - -@tilelang.jit(pass_configs=pass_configs) -def act_quant_kernel( - N, block_size=128, in_dtype=BF16, out_dtype=FP8, scale_dtype=FP32, - round_scale=False, inplace=False -): - """Block-wise FP8 quantization. inplace=True does fused quant+dequant back to BF16.""" - M = T.symbolic("M") - fp8_min = -448.0 - fp8_max = 448.0 - fp8_max_inv = 1 / fp8_max - num_stages = 0 if round_scale or inplace else 2 - blk_m = 32 - group_size = block_size - # Internal computation in FP32; scale_dtype controls output storage format. - compute_dtype = FP32 - out_dtype = in_dtype if inplace else out_dtype - - @T.prim_func - def act_quant_kernel_( - X: T.Tensor[(M, N), in_dtype], - Y: T.Tensor[(M, N), out_dtype], - S: T.Tensor[(M, T.ceildiv(N, group_size)), scale_dtype], - ): - with T.Kernel(T.ceildiv(M, blk_m), T.ceildiv(N, group_size), threads=128) as ( - pid_m, - pid_n, - ): - x_shared = T.alloc_shared((blk_m, group_size), in_dtype) - x_local = T.alloc_fragment((blk_m, group_size), in_dtype) - amax_local = T.alloc_fragment((blk_m,), compute_dtype) - s_local = T.alloc_fragment((blk_m,), compute_dtype) - y_local = T.alloc_fragment((blk_m, group_size), out_dtype) - y_shared = T.alloc_shared((blk_m, group_size), out_dtype) - - for _ in T.Pipelined(1, num_stages=num_stages): - T.copy(X[pid_m * blk_m, pid_n * group_size], x_shared) - T.copy(x_shared, x_local) - T.reduce_absmax(x_local, amax_local, dim=1) - for i in T.Parallel(blk_m): - amax_local[i] = T.max(amax_local[i], 1e-4) - if round_scale: - s_local[i] = fast_round_scale(amax_local[i], fp8_max_inv) - else: - s_local[i] = amax_local[i] * fp8_max_inv - if inplace: - for i, j in T.Parallel(blk_m, group_size): - y_local[i, j] = T.Cast( - out_dtype, - T.Cast(compute_dtype, T.Cast(out_dtype, T.clamp( - x_local[i, j] / s_local[i], fp8_min, fp8_max - ))) * s_local[i], - ) - else: - for i, j in T.Parallel(blk_m, group_size): - y_local[i, j] = T.clamp( - x_local[i, j] / s_local[i], fp8_min, fp8_max - ) - for i in T.Parallel(blk_m): - S[pid_m * blk_m + i, pid_n] = T.Cast(scale_dtype, s_local[i]) - T.copy(y_local, y_shared) - T.copy(y_shared, Y[pid_m * blk_m, pid_n * group_size]) - - return act_quant_kernel_ - - -def act_quant( - x: torch.Tensor, block_size: int = 128, scale_fmt: Optional[str] = None, - scale_dtype: torch.dtype = torch.float32, inplace: bool = False, -) -> torch.Tensor: - """Block-wise FP8 quantization. inplace=True does fused quant+dequant back to BF16. - When scale_fmt is set, scales are rounded to power-of-2 (MXFP).""" - N = x.size(-1) - assert N % block_size == 0 - tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32 - z = x.contiguous() - y = torch.empty_like(z) if inplace else torch.empty_like(z, dtype=torch.float8_e4m3fn) - s = z.new_empty(*z.size()[:-1], N // block_size, dtype=scale_dtype) - kernel = act_quant_kernel( - N, block_size, scale_dtype=tl_dtype, - round_scale=scale_fmt is not None, inplace=inplace, - ) - kernel(z.view(-1, N), y.view(-1, N), s.view(-1, N // block_size)) - if inplace: - x.copy_(y) - return x - return y, s - - -@tilelang.jit(pass_configs=pass_configs) -def fp4_quant_kernel( - N, block_size=32, in_dtype=BF16, scale_dtype=FE8M0, inplace=False -): - """Block-wise FP4 quantization. Power-of-2 scale via bit ops. inplace=True does fused quant+dequant.""" - M = T.symbolic("M") - fp4_max = 6.0 - fp4_max_inv = 1.0 / fp4_max - blk_m = 32 - group_size = block_size - compute_dtype = FP32 - out_dtype = in_dtype if inplace else FP4 - - @T.prim_func - def fp4_quant_kernel_( - X: T.Tensor[(M, N), in_dtype], - Y: T.Tensor[(M, N), out_dtype], - S: T.Tensor[(M, T.ceildiv(N, group_size)), scale_dtype], - ): - with T.Kernel(T.ceildiv(M, blk_m), T.ceildiv(N, group_size), threads=128) as ( - pid_m, - pid_n, - ): - x_shared = T.alloc_shared((blk_m, group_size), in_dtype) - x_local = T.alloc_fragment((blk_m, group_size), in_dtype) - amax_local = T.alloc_fragment((blk_m,), compute_dtype) - s_local = T.alloc_fragment((blk_m,), compute_dtype) - y_local = T.alloc_fragment((blk_m, group_size), out_dtype) - y_shared = T.alloc_shared((blk_m, group_size), out_dtype) - - for _ in T.Pipelined(1, num_stages=2): - T.copy(X[pid_m * blk_m, pid_n * group_size], x_shared) - T.copy(x_shared, x_local) - T.reduce_absmax(x_local, amax_local, dim=1) - for i in T.Parallel(blk_m): - amax_local[i] = T.max(amax_local[i], 6 * (2**-126)) - s_local[i] = fast_round_scale(amax_local[i], fp4_max_inv) - if inplace: - for i, j in T.Parallel(blk_m, group_size): - y_local[i, j] = T.Cast( - out_dtype, - T.Cast(compute_dtype, T.Cast(FP4, T.clamp( - x_local[i, j] / s_local[i], -fp4_max, fp4_max - ))) * s_local[i], - ) - else: - for i, j in T.Parallel(blk_m, group_size): - y_local[i, j] = T.clamp( - x_local[i, j] / s_local[i], -fp4_max, fp4_max - ) - for i in T.Parallel(blk_m): - S[pid_m * blk_m + i, pid_n] = T.Cast(scale_dtype, s_local[i]) - T.copy(y_local, y_shared) - T.copy(y_shared, Y[pid_m * blk_m, pid_n * group_size]) - - return fp4_quant_kernel_ - - -def fp4_act_quant( - x: torch.Tensor, block_size: int = 32, inplace: bool = False, -) -> torch.Tensor: - """Block-wise FP4 quantization. inplace=True does fused quant+dequant back to BF16.""" - N = x.size(-1) - assert N % block_size == 0 - z = x.contiguous() - y = torch.empty_like(z) if inplace else z.new_empty(*z.shape[:-1], N // 2, dtype=torch.float4_e2m1fn_x2) - s = z.new_empty(*z.size()[:-1], N // block_size, dtype=torch.float8_e8m0fnu) - kernel = fp4_quant_kernel(N, block_size, inplace=inplace) - kernel(z.view(-1, N), y.view(-1, y.size(-1)), s.view(-1, N // block_size)) - if inplace: - x.copy_(y) - return x - return y, s - - -@tilelang.jit(pass_configs=pass_configs) -def fp8_gemm_kernel(N, K, out_dtype=BF16, accum_dtype=FP32, scale_dtype=FP32): - assert out_dtype in [BF16, FP32] - - M = T.symbolic("M") - group_size = 128 - block_M = 32 - block_N = 128 - block_K = 128 - - @T.prim_func - def fp8_gemm_kernel_( - A: T.Tensor[(M, K), FP8], - B: T.Tensor[(N, K), FP8], - C: T.Tensor[(M, N), out_dtype], - scales_a: T.Tensor[(M, T.ceildiv(K, group_size)), scale_dtype], - scales_b: T.Tensor[(T.ceildiv(N, group_size), T.ceildiv(K, group_size)), scale_dtype], - ): - with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as ( - bx, - by, - ): - A_shared = T.alloc_shared((block_M, block_K), FP8) - B_shared = T.alloc_shared((block_N, block_K), FP8) - C_shared = T.alloc_shared((block_M, block_N), out_dtype) - Scale_C_shared = T.alloc_shared((block_M), FP32) - C_local = T.alloc_fragment((block_M, block_N), accum_dtype) - C_local_accum = T.alloc_fragment((block_M, block_N), accum_dtype) - - # Improve L2 Cache - T.use_swizzle(panel_size=10) - T.clear(C_local) - T.clear(C_local_accum) - - K_iters = T.ceildiv(K, block_K) - for k in T.Pipelined(K_iters, num_stages=4): - T.copy(A[by * block_M, k * block_K], A_shared) - T.copy(B[bx * block_N, k * block_K], B_shared) - # Cast scales to FP32 for computation; scales_b has one value per block_N group - Scale_B = T.Cast(FP32, scales_b[bx * block_N // group_size, k]) - for i in T.Parallel(block_M): - Scale_C_shared[i] = T.Cast(FP32, scales_a[by * block_M + i, k]) * Scale_B - - T.gemm(A_shared, B_shared, C_local, transpose_B=True) - # Separate accumulator for scale-corrected results (2x accumulation precision) - for i, j in T.Parallel(block_M, block_N): - C_local_accum[i, j] += C_local[i, j] * Scale_C_shared[i] - T.clear(C_local) - T.copy(C_local_accum, C_shared) - T.copy(C_shared, C[by * block_M, bx * block_N]) - - return fp8_gemm_kernel_ - - -def fp8_gemm( - a: torch.Tensor, a_s: torch.Tensor, b: torch.Tensor, b_s: torch.Tensor, - scale_dtype: torch.dtype = torch.float32, -) -> torch.Tensor: - """C[M,N] = A[M,K] @ B[N,K]^T with per-128 block FP8 scaling on both A and B.""" - assert a.is_contiguous() and b.is_contiguous(), "Input tensors must be contiguous" - assert a_s.is_contiguous() and b_s.is_contiguous(), ( - "Scaling factor tensors must be contiguous" - ) - tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32 - K = a.size(-1) - M = a.numel() // K - N = b.size(0) - c = a.new_empty(*a.size()[:-1], N, dtype=torch.get_default_dtype()) - kernel = fp8_gemm_kernel(N, K, scale_dtype=tl_dtype) - kernel(a.view(M, K), b, c.view(M, N), a_s.view(M, -1), b_s) - return c - - -@tilelang.jit(pass_configs=pass_configs) -def sparse_attn_kernel(h: int, d: int, scale=None): - """Sparse multi-head attention via index gathering + online softmax (FlashAttention-style). - For each (batch, seq_pos), gathers top-k KV positions by index, computes attention - with numerically stable running max/sum, and includes a learnable attn_sink bias.""" - b = T.symbolic("b") - m = T.symbolic("m") - n = T.symbolic("n") - topk = T.symbolic("topk") - if scale is None: - scale = (1.0 / d) ** 0.5 - - num_stages = 2 - threads = 256 - block = 64 - num_blocks = tilelang.cdiv(topk, block) - - @T.prim_func - def sparse_attn_kernel_( - q: T.Tensor[(b, m, h, d), BF16], - kv: T.Tensor[(b, n, d), BF16], - o: T.Tensor[(b, m, h, d), BF16], - attn_sink: T.Tensor[(h,), FP32], - topk_idxs: T.Tensor[(b, m, topk), INT32], - ): - with T.Kernel(m, b, threads=threads) as (bx, by): - q_shared = T.alloc_shared((h, d), BF16) - kv_shared = T.alloc_shared((block, d), BF16) - o_shared = T.alloc_shared((h, d), BF16) - acc_s_cast = T.alloc_shared((h, block), BF16) - - idxs = T.alloc_fragment(block, INT32) - acc_s = T.alloc_fragment((h, block), FP32) - acc_o = T.alloc_fragment((h, d), FP32) - scores_max = T.alloc_fragment(h, FP32) - scores_max_prev = T.alloc_fragment(h, FP32) - scores_scale = T.alloc_fragment(h, FP32) - scores_sum = T.alloc_fragment(h, FP32) - sum_exp = T.alloc_fragment(h, FP32) - - T.clear(acc_o) - T.clear(sum_exp) - T.fill(scores_max, -T.infinity(FP32)) - T.copy(q[by, bx, :, :], q_shared) - - for t in T.Pipelined(num_blocks, num_stages=num_stages): - for i in T.Parallel(block): - idxs[i] = T.if_then_else(t * block + i < topk, topk_idxs[by, bx, t * block + i], -1) - for i, j in T.Parallel(block, d): - kv_shared[i, j] = T.if_then_else(idxs[i] != -1, kv[by, idxs[i], j], 0) - for i, j in T.Parallel(h, block): - acc_s[i, j] = T.if_then_else(idxs[j] != -1, 0, -T.infinity(FP32)) - T.gemm(q_shared, kv_shared, acc_s, transpose_B=True, policy=T.GemmWarpPolicy.FullRow) - for i, j in T.Parallel(h, block): - acc_s[i, j] *= scale - T.copy(scores_max, scores_max_prev) - T.reduce_max(acc_s, scores_max, dim=1, clear=False) - for i in T.Parallel(h): - scores_scale[i] = T.exp(scores_max_prev[i] - scores_max[i]) - for i, j in T.Parallel(h, block): - acc_s[i, j] = T.exp(acc_s[i, j] - scores_max[i]) - T.reduce_sum(acc_s, scores_sum, dim=1) - for i in T.Parallel(h): - sum_exp[i] = sum_exp[i] * scores_scale[i] + scores_sum[i] - T.copy(acc_s, acc_s_cast) - for i, j in T.Parallel(h, d): - acc_o[i, j] *= scores_scale[i] - T.gemm(acc_s_cast, kv_shared, acc_o, policy=T.GemmWarpPolicy.FullRow) - - for i in T.Parallel(h): - sum_exp[i] += T.exp(attn_sink[i] - scores_max[i]) - for i, j in T.Parallel(h, d): - acc_o[i, j] /= sum_exp[i] - T.copy(acc_o, o_shared) - T.copy(o_shared, o[by, bx, :, :]) - - return sparse_attn_kernel_ - - -def sparse_attn( - q: torch.Tensor, kv: torch.Tensor, attn_sink: torch.Tensor, topk_idxs: torch.Tensor, softmax_scale: float -) -> torch.Tensor: - b, s, h, d = q.size() - # Pad heads to 16 for kernel efficiency (stripped after) - if h < 16: - q = torch.cat([q, q.new_zeros(b, s, 16 - h, d)], dim=2) - attn_sink = torch.cat([attn_sink, attn_sink.new_zeros(16 - h)]) - o = torch.empty_like(q) - kernel = sparse_attn_kernel(q.size(2), d, softmax_scale) - kernel(q, kv, o, attn_sink, topk_idxs) - if h < 16: - o = o.narrow(2, 0, h).contiguous() - return o - - -@tilelang.jit(pass_configs=pass_configs) -def hc_split_sinkhorn_kernel(hc: int, sinkhorn_iters: int, eps: float): - n = T.symbolic("n") - mix_hc = (2 + hc) * hc - threads = 64 - - @T.prim_func - def hc_split_sinkhorn_kernel_( - mixes: T.Tensor[(n, mix_hc), FP32], - hc_scale: T.Tensor[(3,), FP32], - hc_base: T.Tensor[(mix_hc,), FP32], - pre: T.Tensor[(n, hc), FP32], - post: T.Tensor[(n, hc), FP32], - comb: T.Tensor[(n, hc, hc), FP32], - ): - with T.Kernel(n, threads=threads) as i: - mixes_shared = T.alloc_shared(mix_hc, FP32) - comb_frag = T.alloc_fragment((hc, hc), FP32) - T.copy(mixes[i, :], mixes_shared) - - for j in T.Parallel(hc): - pre[i, j] = T.sigmoid(mixes_shared[j] * hc_scale[0] + hc_base[j]) + eps - for j in T.Parallel(hc): - post[i, j] = 2 * T.sigmoid(mixes_shared[j + hc] * hc_scale[1] + hc_base[j + hc]) - for j, k in T.Parallel(hc, hc): - comb_frag[j, k] = mixes_shared[j * hc + k + hc * 2] * hc_scale[2] + hc_base[j * hc + k + hc * 2] - - row_sum = T.alloc_fragment(hc, FP32) - col_sum = T.alloc_fragment(hc, FP32) - - # comb = comb.softmax(-1) + eps - row_max = T.alloc_fragment(hc, FP32) - T.reduce_max(comb_frag, row_max, dim=1) - for j, k in T.Parallel(hc, hc): - comb_frag[j, k] = T.exp(comb_frag[j, k] - row_max[j]) - T.reduce_sum(comb_frag, row_sum, dim=1) - for j, k in T.Parallel(hc, hc): - comb_frag[j, k] = comb_frag[j, k] / row_sum[j] + eps - - # comb = comb / (comb.sum(-2) + eps) - T.reduce_sum(comb_frag, col_sum, dim=0) - for j, k in T.Parallel(hc, hc): - comb_frag[j, k] = comb_frag[j, k] / (col_sum[k] + eps) - - for _ in T.serial(sinkhorn_iters - 1): - # comb = comb / (comb.sum(-1) + eps) - T.reduce_sum(comb_frag, row_sum, dim=1) - for j, k in T.Parallel(hc, hc): - comb_frag[j, k] = comb_frag[j, k] / (row_sum[j] + eps) - # comb = comb / (comb.sum(-2) + eps) - T.reduce_sum(comb_frag, col_sum, dim=0) - for j, k in T.Parallel(hc, hc): - comb_frag[j, k] = comb_frag[j, k] / (col_sum[k] + eps) - - T.copy(comb_frag, comb[i, :, :]) - - return hc_split_sinkhorn_kernel_ - - -def hc_split_sinkhorn(mixes: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor, hc_mult: int = 4, sinkhorn_iters: int = 20, eps: float = 1e-6): - b, s, _ = mixes.size() - pre = mixes.new_empty(b, s, hc_mult) - post = mixes.new_empty(b, s, hc_mult) - comb = mixes.new_empty(b, s, hc_mult, hc_mult) - kernel = hc_split_sinkhorn_kernel(hc_mult, sinkhorn_iters, eps) - kernel(mixes.view(-1, (2 + hc_mult) * hc_mult), hc_scale, hc_base, - pre.view(-1, hc_mult), post.view(-1, hc_mult), comb.view(-1, hc_mult, hc_mult)) - return pre, post, comb - - -@tilelang.jit(pass_configs=pass_configs) -def fp4_gemm_kernel(N, K, out_dtype=BF16, accum_dtype=FP32, scale_dtype=FP32): - """FP8 act x FP4 weight GEMM kernel. - - C[M, N] = A_fp8[M, K] @ B_fp4[N, K]^T - - Act: 1x128 quant on K (reduce dim), FP8 with configurable scale dtype - Weight: 1x32 quant on K (reduce dim), FP4 with E8M0 scale - - B is stored as [N, K//2] in float4_e2m1fn_x2, logical [N, K] in fp4. - The FP4 values are packed along the K (last) dimension. - - Strategy: load FP4 sub-blocks of size [block_N, sub_K] (sub_K=32), - cast FP4 to FP8 via float, then do FP8xFP8 GEMM. - Apply act scale (per 128 on K) and weight scale (per 32 on K) to the accumulator. - """ - M = T.symbolic("M") - act_group_size = 128 - weight_group_size = 32 - block_M = 32 - block_N = 128 - block_K = 32 # matches weight_group_size for simple scale handling - n_sub = act_group_size // block_K # 4 sub-blocks per act scale group - - @T.prim_func - def fp4_gemm_kernel_( - A: T.Tensor[(M, K), FP8], - B: T.Tensor[(N, K), FP4], - C: T.Tensor[(M, N), out_dtype], - scales_a: T.Tensor[(M, T.ceildiv(K, act_group_size)), scale_dtype], - scales_b: T.Tensor[(N, T.ceildiv(K, weight_group_size)), scale_dtype], - ): - with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as ( - bx, - by, - ): - A_shared = T.alloc_shared((block_M, block_K), FP8) - B_fp4_shared = T.alloc_shared((block_N, block_K), FP4) - B_shared = T.alloc_shared((block_N, block_K), FP8) - C_shared = T.alloc_shared((block_M, block_N), out_dtype) - C_local = T.alloc_fragment((block_M, block_N), accum_dtype) - C_local_accum = T.alloc_fragment((block_M, block_N), accum_dtype) - scale_a_frag = T.alloc_fragment((block_M,), FP32) - scale_b_frag = T.alloc_fragment((block_N,), FP32) - - T.use_swizzle(panel_size=10) - T.clear(C_local) - T.clear(C_local_accum) - - K_iters = T.ceildiv(K, block_K) - for k in T.Pipelined(K_iters, num_stages=2): - T.copy(A[by * block_M, k * block_K], A_shared) - T.copy(B[bx * block_N, k * block_K], B_fp4_shared) - # FP4->FP8 cast must go through FP32 to avoid ambiguous C++ overload - for i, j in T.Parallel(block_N, block_K): - B_shared[i, j] = T.Cast(FP8, T.Cast(FP32, B_fp4_shared[i, j])) - - # Weight scale: per 32 on K, indexed by k (each k is one block_K=32) - for i in T.Parallel(block_N): - scale_b_frag[i] = T.Cast(FP32, scales_b[bx * block_N + i, k]) - - # Act scale: per 128 on K, indexed by k // 4 - for i in T.Parallel(block_M): - scale_a_frag[i] = T.Cast(FP32, scales_a[by * block_M + i, k // n_sub]) - - T.gemm(A_shared, B_shared, C_local, transpose_B=True) - - for i, j in T.Parallel(block_M, block_N): - C_local_accum[i, j] += C_local[i, j] * scale_a_frag[i] * scale_b_frag[j] - T.clear(C_local) - - T.copy(C_local_accum, C_shared) - T.copy(C_shared, C[by * block_M, bx * block_N]) - - return fp4_gemm_kernel_ - - -def fp4_gemm( - a: torch.Tensor, a_s: torch.Tensor, b: torch.Tensor, b_s: torch.Tensor, - scale_dtype: torch.dtype = torch.float32, -) -> torch.Tensor: - """C[M,N] = A_fp8[M,K] @ B_fp4[N,K]^T. - A has per-128 act scale; B has per-32 E8M0 weight scale. - B is stored as [N, K//2] in float4_e2m1fn_x2 (2 FP4 values per byte, packed along K).""" - assert a.is_contiguous() and b.is_contiguous(), "Input tensors must be contiguous" - assert a_s.is_contiguous() and b_s.is_contiguous(), ( - "Scaling factor tensors must be contiguous" - ) - tl_dtype = FE8M0 if scale_dtype == torch.float8_e8m0fnu else FP32 - K = a.size(-1) - M = a.numel() // K - N = b.size(0) - c = a.new_empty(*a.size()[:-1], N, dtype=torch.get_default_dtype()) - kernel = fp4_gemm_kernel(N, K, scale_dtype=tl_dtype) - kernel(a.view(M, K), b, c.view(M, N), a_s.view(M, -1), b_s) - return c diff --git a/TEMP/dsv4thing/model.py b/TEMP/dsv4thing/model.py deleted file mode 100644 index 167ade8f..00000000 --- a/TEMP/dsv4thing/model.py +++ /dev/null @@ -1,827 +0,0 @@ -import math -from dataclasses import dataclass -from typing import Tuple, Optional, Literal -from functools import lru_cache -from contextlib import contextmanager - -import torch -from torch import nn -import torch.nn.functional as F -import torch.distributed as dist - -from kernel import act_quant, fp4_act_quant, fp8_gemm, fp4_gemm, sparse_attn, hc_split_sinkhorn - - -world_size = 1 -rank = 0 -block_size = 128 -fp4_block_size = 32 -default_dtype = torch.bfloat16 -scale_fmt = None -scale_dtype = torch.float32 - - -@contextmanager -def set_dtype(dtype): - """Temporarily override torch default dtype, restoring it on exit (even if an exception occurs).""" - prev = torch.get_default_dtype() - torch.set_default_dtype(dtype) - try: - yield - finally: - torch.set_default_dtype(prev) - -@dataclass -class ModelArgs: - """Model hyperparameters. Field names match the config JSON keys.""" - max_batch_size: int = 4 - max_seq_len: int = 4096 - dtype: Literal["bf16", "fp8"] = "fp8" - scale_fmt: Literal[None, "ue8m0"] = "ue8m0" - expert_dtype: Literal[None, "fp4"] = None - scale_dtype: Literal["fp32", "fp8"] = "fp8" - vocab_size: int = 129280 - dim: int = 4096 - moe_inter_dim: int = 4096 - n_layers: int = 7 - n_hash_layers: int = 0 - n_mtp_layers: int = 1 - n_heads: int = 64 - # moe - n_routed_experts: int = 8 - n_shared_experts: int = 1 - n_activated_experts: int = 2 - score_func: Literal["softmax", "sigmoid", "sqrtsoftplus"] = "sqrtsoftplus" - route_scale: float = 1. - swiglu_limit: float = 0. - # mqa - q_lora_rank: int = 1024 - head_dim: int = 512 - rope_head_dim: int = 64 - norm_eps: float = 1e-6 - o_groups: int = 8 - o_lora_rank: int = 1024 - window_size: int = 128 - compress_ratios: Tuple[int] = (0, 0, 4, 128, 4, 128, 4, 0) - # yarn - compress_rope_theta: float = 40000.0 - original_seq_len: int = 0 - rope_theta: float = 10000.0 - rope_factor: float = 40 - beta_fast: int = 32 - beta_slow: int = 1 - # index - index_n_heads: int = 64 - index_head_dim: int = 128 - index_topk: int = 512 - # hc - hc_mult: int = 4 - hc_sinkhorn_iters: int = 20 - hc_eps: float = 1e-6 - - -class ParallelEmbedding(nn.Module): - """Embedding sharded along the vocab dimension. Each rank holds vocab_size // world_size rows. - Out-of-range indices are zero-masked before all_reduce to combine partial embeddings.""" - def __init__(self, vocab_size: int, dim: int): - super().__init__() - self.vocab_size = vocab_size - self.dim = dim - assert vocab_size % world_size == 0, f"Vocabulary size must be divisible by world size (world_size={world_size})" - self.part_vocab_size = (vocab_size // world_size) - self.vocab_start_idx = rank * self.part_vocab_size - self.vocab_end_idx = self.vocab_start_idx + self.part_vocab_size - self.weight = nn.Parameter(torch.empty(self.part_vocab_size, self.dim)) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if world_size > 1: - mask = (x < self.vocab_start_idx) | (x >= self.vocab_end_idx) - x = x - self.vocab_start_idx - x[mask] = 0 - y = F.embedding(x, self.weight) - if world_size > 1: - y[mask] = 0 - dist.all_reduce(y) - return y - - -def linear(x: torch.Tensor, weight: torch.Tensor, bias: Optional[torch.Tensor] = None) -> torch.Tensor: - """Dispatches to fp4_gemm / fp8_gemm / F.linear based on weight dtype. - For quantized weights, x is first quantized to FP8 via act_quant.""" - assert bias is None - - if weight.dtype == torch.float4_e2m1fn_x2: - x, s = act_quant(x, block_size, scale_fmt, scale_dtype) - return fp4_gemm(x, s, weight, weight.scale, scale_dtype) - elif weight.dtype == torch.float8_e4m3fn: - x, s = act_quant(x, block_size, scale_fmt, scale_dtype) - return fp8_gemm(x, s, weight, weight.scale, scale_dtype) - else: - return F.linear(x, weight) - - -class Linear(nn.Module): - """Linear layer supporting BF16, FP8, and FP4 weight formats with per-block scaling.""" - - def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None): - super().__init__() - self.in_features = in_features - self.out_features = out_features - dtype = dtype or default_dtype - if dtype == torch.float4_e2m1fn_x2: - # FP4: weight is [out, in//2] in float4_e2m1fn_x2, logically [out, in] in fp4 - # Scale is [out, in//32] in float8_e8m0fnu (1 scale per 32 fp4 elements along K) - self.weight = nn.Parameter(torch.empty(out_features, in_features // 2, dtype=torch.float4_e2m1fn_x2)) - scale_out_features = out_features - scale_in_features = in_features // fp4_block_size - self.weight.scale = self.scale = nn.Parameter(torch.empty(scale_out_features, scale_in_features, dtype=torch.float8_e8m0fnu)) - elif dtype == torch.float8_e4m3fn: - self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype)) - scale_out_features = (out_features + block_size - 1) // block_size - scale_in_features = (in_features + block_size - 1) // block_size - self.weight.scale = self.scale = nn.Parameter(torch.empty(scale_out_features, scale_in_features, dtype=torch.float8_e8m0fnu)) - else: - self.weight = nn.Parameter(torch.empty(out_features, in_features, dtype=dtype)) - self.register_parameter("scale", None) - if bias: - self.bias = nn.Parameter(torch.empty(out_features)) - else: - self.register_parameter("bias", None) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return linear(x, self.weight, self.bias) - - -class ColumnParallelLinear(Linear): - """Shards output dim across TP ranks. No all-reduce needed on output.""" - def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None): - assert out_features % world_size == 0, f"Output features must be divisible by world size (world_size={world_size})" - self.part_out_features = out_features // world_size - super().__init__(in_features, self.part_out_features, bias, dtype) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return linear(x, self.weight, self.bias) - - -class RowParallelLinear(Linear): - """Shards input dim across TP ranks. All-reduce on output to sum partial results.""" - def __init__(self, in_features: int, out_features: int, bias: bool = False, dtype = None): - assert in_features % world_size == 0, f"Input features must be divisible by world size (world_size={world_size})" - self.part_in_features = in_features // world_size - super().__init__(self.part_in_features, out_features, bias, dtype) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - y = linear(x, self.weight, None) - if world_size > 1: - y = y.float() - dist.all_reduce(y) - if self.bias is not None: - y += self.bias - return y.type_as(x) - - -class RMSNorm(nn.Module): - def __init__(self, dim: int, eps: float = 1e-6): - super().__init__() - self.dim = dim - self.eps = eps - # rmsnorm in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for convenient. - self.weight = nn.Parameter(torch.ones(dim, dtype=torch.float32)) - - def forward(self, x: torch.Tensor): - dtype = x.dtype - x = x.float() - var = x.square().mean(-1, keepdim=True) - x = x * torch.rsqrt(var + self.eps) - return (self.weight * x).to(dtype) - - -@lru_cache(2) -def precompute_freqs_cis(dim, seqlen, original_seq_len, base, factor, beta_fast, beta_slow) -> torch.Tensor: - """Precomputes complex exponentials for rotary embeddings with YaRN scaling. - When original_seq_len > 0, applies frequency interpolation with a smooth - linear ramp between beta_fast and beta_slow correction ranges.""" - - def find_correction_dim(num_rotations, dim, base, max_seq_len): - return dim * math.log(max_seq_len / (num_rotations * 2 * math.pi)) / (2 * math.log(base)) - - def find_correction_range(low_rot, high_rot, dim, base, max_seq_len): - low = math.floor(find_correction_dim(low_rot, dim, base, max_seq_len)) - high = math.ceil(find_correction_dim(high_rot, dim, base, max_seq_len)) - return max(low, 0), min(high, dim-1) - - def linear_ramp_factor(min, max, dim): - if min == max: - max += 0.001 - linear_func = (torch.arange(dim, dtype=torch.float32) - min) / (max - min) - ramp_func = torch.clamp(linear_func, 0, 1) - return ramp_func - - freqs = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) - if original_seq_len > 0: - low, high = find_correction_range(beta_fast, beta_slow, dim, base, original_seq_len) - smooth = 1 - linear_ramp_factor(low, high, dim // 2) - freqs = freqs / factor * (1 - smooth) + freqs * smooth - - t = torch.arange(seqlen) - freqs = torch.outer(t, freqs) - freqs_cis = torch.polar(torch.ones_like(freqs), freqs) - return freqs_cis - - -def apply_rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor, inverse: bool = False) -> torch.Tensor: - """Applies rotary positional embeddings in-place. Uses conjugate for inverse (de-rotation).""" - y = x - x = torch.view_as_complex(x.float().unflatten(-1, (-1, 2))) - if inverse: - freqs_cis = freqs_cis.conj() - if x.ndim == 3: - freqs_cis = freqs_cis.view(1, x.size(1), x.size(-1)) - else: - freqs_cis = freqs_cis.view(1, x.size(1), 1, x.size(-1)) - x = torch.view_as_real(x * freqs_cis).flatten(-2) - y.copy_(x) - return y - - -def rotate_activation(x: torch.Tensor) -> torch.Tensor: - """Applies randomized Hadamard rotation to spread information across dims before FP8 quant.""" - assert x.dtype == torch.bfloat16 - from fast_hadamard_transform import hadamard_transform - return hadamard_transform(x, scale=x.size(-1) ** -0.5) - - -@lru_cache(1) -def get_window_topk_idxs(window_size: int, bsz: int, seqlen: int, start_pos: int): - if start_pos >= window_size - 1: - start_pos %= window_size - matrix = torch.cat([torch.arange(start_pos + 1, window_size), torch.arange(0, start_pos + 1)], dim=0) - elif start_pos > 0: - matrix = F.pad(torch.arange(start_pos + 1), (0, window_size - start_pos - 1), value=-1) - else: - base = torch.arange(seqlen).unsqueeze(1) - matrix = (base - window_size + 1).clamp(0) + torch.arange(min(seqlen, window_size)) - matrix = torch.where(matrix > base, -1, matrix) - return matrix.unsqueeze(0).expand(bsz, -1, -1) - - -@lru_cache(2) -def get_compress_topk_idxs(ratio: int, bsz: int, seqlen: int, start_pos: int, offset: int): - if start_pos > 0: - matrix = torch.arange(0, (start_pos + 1) // ratio) + offset - else: - matrix = torch.arange(seqlen // ratio).repeat(seqlen, 1) - mask = matrix >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio - matrix = torch.where(mask, -1, matrix + offset) - return matrix.unsqueeze(0).expand(bsz, -1, -1) - - -class Compressor(nn.Module): - """Compresses KV cache via learned gated pooling over `compress_ratio` consecutive tokens. - When overlap=True (ratio==4), uses overlapping windows for smoother compression boundaries.""" - - def __init__(self, args: ModelArgs, compress_ratio: int = 4, head_dim: int = 512, rotate: bool = False): - super().__init__() - self.dim = args.dim - self.head_dim = head_dim - self.rope_head_dim = args.rope_head_dim - self.nope_head_dim = head_dim - args.rope_head_dim - self.compress_ratio = compress_ratio - self.overlap = compress_ratio == 4 - self.rotate = rotate - coff = 1 + self.overlap - - self.ape = nn.Parameter(torch.empty(compress_ratio, coff * self.head_dim, dtype=torch.float32)) - # wkv and wgate in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for convenient. - # When overlap, the first half of dims is for overlapping compression, second half for normal. - self.wkv = Linear(self.dim, coff * self.head_dim, dtype=torch.float32) - self.wgate = Linear(self.dim, coff * self.head_dim, dtype=torch.float32) - self.norm = RMSNorm(self.head_dim, args.norm_eps) - self.kv_cache: torch.Tensor = None # assigned lazily from Attention.kv_cache - # State buffers for decode-phase incremental compression. - # With overlap: state[:, :ratio] = overlapping window, state[:, ratio:] = current window. - self.register_buffer("kv_state", torch.zeros(args.max_batch_size, coff * compress_ratio, coff * self.head_dim, dtype=torch.float32), persistent=False) - self.register_buffer("score_state", torch.full((args.max_batch_size, coff * compress_ratio, coff * self.head_dim), float("-inf"), dtype=torch.float32), persistent=False) - self.freqs_cis: torch.Tensor = None - - def overlap_transform(self, tensor: torch.Tensor, value=0): - # tensor: [b,s,r,2d] - b, s, _, _ = tensor.size() - ratio, d = self.compress_ratio, self.head_dim - new_tensor = tensor.new_full((b, s, 2 * ratio, d), value) - new_tensor[:, :, ratio:] = tensor[:, :, :, d:] - new_tensor[:, 1:, :ratio] = tensor[:, :-1, :, :d] - return new_tensor - - def forward(self, x: torch.Tensor, start_pos: int): - assert self.kv_cache is not None - bsz, seqlen, _ = x.size() - ratio, overlap, d, rd = self.compress_ratio, self.overlap, self.head_dim, self.rope_head_dim - dtype = x.dtype - # compression need fp32 - x = x.float() - kv = self.wkv(x) - score = self.wgate(x) - if start_pos == 0: - should_compress = seqlen >= ratio - remainder = seqlen % ratio - cutoff = seqlen - remainder - offset = ratio if overlap else 0 - if overlap and cutoff >= ratio: - self.kv_state[:bsz, :ratio] = kv[:, cutoff-ratio : cutoff] - self.score_state[:bsz, :ratio] = score[:, cutoff-ratio : cutoff] + self.ape - if remainder > 0: - kv, self.kv_state[:bsz, offset : offset+remainder] = kv.split([cutoff, remainder], dim=1) - self.score_state[:bsz, offset : offset+remainder] = score[:, cutoff:] + self.ape[:remainder] - score = score[:, :cutoff] - kv = kv.unflatten(1, (-1, ratio)) - score = score.unflatten(1, (-1, ratio)) + self.ape - if overlap: - kv = self.overlap_transform(kv, 0) - score = self.overlap_transform(score, float("-inf")) - kv = (kv * score.softmax(dim=2)).sum(dim=2) - else: - should_compress = (start_pos + 1) % self.compress_ratio == 0 - score += self.ape[start_pos % ratio] - if overlap: - self.kv_state[:bsz, ratio + start_pos % ratio] = kv.squeeze(1) - self.score_state[:bsz, ratio + start_pos % ratio] = score.squeeze(1) - if should_compress: - kv_state = torch.cat([self.kv_state[:bsz, :ratio, :d], self.kv_state[:bsz, ratio:, d:]], dim=1) - score_state = torch.cat([self.score_state[:bsz, :ratio, :d], self.score_state[:bsz, ratio:, d:]], dim=1) - kv = (kv_state * score_state.softmax(dim=1)).sum(dim=1, keepdim=True) - self.kv_state[:bsz, :ratio] = self.kv_state[:bsz, ratio:] - self.score_state[:bsz, :ratio] = self.score_state[:bsz, ratio:] - else: - self.kv_state[:bsz, start_pos % ratio] = kv.squeeze(1) - self.score_state[:bsz, start_pos % ratio] = score.squeeze(1) - if should_compress: - kv = (self.kv_state[:bsz] * self.score_state[:bsz].softmax(dim=1)).sum(dim=1, keepdim=True) - if not should_compress: - return - kv = self.norm(kv.to(dtype)) - if start_pos == 0: - freqs_cis = self.freqs_cis[:cutoff:ratio] - else: - freqs_cis = self.freqs_cis[start_pos + 1 - self.compress_ratio].unsqueeze(0) - apply_rotary_emb(kv[..., -rd:], freqs_cis) - if self.rotate: - kv = rotate_activation(kv) - fp4_act_quant(kv, fp4_block_size, True) - else: - act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True) - if start_pos == 0: - self.kv_cache[:bsz, :seqlen // ratio] = kv - else: - self.kv_cache[:bsz, start_pos // ratio] = kv.squeeze(1) - return kv - - -class Indexer(torch.nn.Module): - """Selects top-k compressed KV positions for sparse attention via learned scoring. - Has its own Compressor (with Hadamard rotation) to build compressed KV for scoring.""" - - def __init__(self, args: ModelArgs, compress_ratio: int = 4): - super().__init__() - self.dim = args.dim - self.n_heads = args.index_n_heads - self.n_local_heads = args.index_n_heads // world_size - self.head_dim = args.index_head_dim - self.rope_head_dim = args.rope_head_dim - self.index_topk = args.index_topk - self.q_lora_rank = args.q_lora_rank - self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.head_dim) - self.weights_proj = ColumnParallelLinear(self.dim, self.n_heads, dtype=torch.bfloat16) - self.softmax_scale = self.head_dim ** -0.5 - self.compress_ratio = compress_ratio - - self.compressor = Compressor(args, compress_ratio, self.head_dim, True) - self.register_buffer("kv_cache", torch.zeros(args.max_batch_size, args.max_seq_len // compress_ratio, self.head_dim), persistent=False) - self.freqs_cis = None - - def forward(self, x: torch.Tensor, qr: torch.Tensor, start_pos: int, offset: int): - bsz, seqlen, _ = x.size() - freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen] - ratio = self.compress_ratio - rd = self.rope_head_dim - end_pos = start_pos + seqlen - if self.compressor.kv_cache is None: - self.compressor.kv_cache = self.kv_cache - self.compressor.freqs_cis = self.freqs_cis - q = self.wq_b(qr) - q = q.unflatten(-1, (self.n_local_heads, self.head_dim)) - apply_rotary_emb(q[..., -rd:], freqs_cis) - q = rotate_activation(q) - # use fp4 simulation for q and kv in indexer - fp4_act_quant(q, fp4_block_size, True) - self.compressor(x, start_pos) - weights = self.weights_proj(x) * (self.softmax_scale * self.n_heads ** -0.5) - # We performed QAT here, kv could also use fp8 format, though current implementation uses bf16 - index_score = torch.einsum("bshd,btd->bsht", q, self.kv_cache[:bsz, :end_pos // ratio]) - index_score = (index_score.relu_() * weights.unsqueeze(-1)).sum(dim=2) - if world_size > 1: - dist.all_reduce(index_score) - if start_pos == 0: - mask = torch.arange(seqlen // ratio).repeat(seqlen, 1) >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio - index_score += torch.where(mask, float("-inf"), 0) - topk_idxs = index_score.topk(min(self.index_topk, end_pos // ratio), dim=-1)[1] - if start_pos == 0: - mask = topk_idxs >= torch.arange(1, seqlen + 1).unsqueeze(1) // ratio - topk_idxs = torch.where(mask, -1, topk_idxs + offset) - else: - topk_idxs += offset - return topk_idxs - - -class Attention(nn.Module): - """Multi-head Latent Attention (MLA) with sliding window + optional KV compression. - Uses low-rank Q projection (wq_a -> q_norm -> wq_b) and grouped low-rank O projection.""" - def __init__(self, layer_id: int, args: ModelArgs): - super().__init__() - self.layer_id = layer_id - self.dim = args.dim - self.n_heads = args.n_heads - self.n_local_heads = args.n_heads // world_size - self.q_lora_rank = args.q_lora_rank - self.o_lora_rank = args.o_lora_rank - self.head_dim = args.head_dim - self.rope_head_dim = args.rope_head_dim - self.nope_head_dim = args.head_dim - args.rope_head_dim - self.n_groups = args.o_groups - self.n_local_groups = self.n_groups // world_size - self.window_size = args.window_size - self.compress_ratio = args.compress_ratios[layer_id] - self.eps = args.norm_eps - - self.attn_sink = nn.Parameter(torch.empty(self.n_local_heads, dtype=torch.float32)) - self.wq_a = Linear(self.dim, self.q_lora_rank) - self.q_norm = RMSNorm(self.q_lora_rank, self.eps) - self.wq_b = ColumnParallelLinear(self.q_lora_rank, self.n_heads * self.head_dim) - self.wkv = Linear(self.dim, self.head_dim) - self.kv_norm = RMSNorm(self.head_dim, self.eps) - self.wo_a = ColumnParallelLinear(self.n_heads * self.head_dim // self.n_groups, self.n_groups * args.o_lora_rank, dtype=torch.bfloat16) - self.wo_b = RowParallelLinear(self.n_groups * args.o_lora_rank, self.dim) - self.softmax_scale = self.head_dim ** -0.5 - - if self.compress_ratio: - self.compressor = Compressor(args, self.compress_ratio, self.head_dim) - if self.compress_ratio == 4: - self.indexer = Indexer(args, self.compress_ratio) - else: - self.indexer = None - - kv_cache_size = args.window_size + (args.max_seq_len // self.compress_ratio if self.compress_ratio else 0) - self.register_buffer("kv_cache", torch.zeros(args.max_batch_size, kv_cache_size, self.head_dim), persistent=False) - if self.compress_ratio: - original_seq_len, rope_theta = args.original_seq_len, args.compress_rope_theta - else: - # disable YaRN and use base rope_theta in pure sliding-window attention - original_seq_len, rope_theta = 0, args.rope_theta - freqs_cis = precompute_freqs_cis(self.rope_head_dim, args.max_seq_len, original_seq_len, - rope_theta, args.rope_factor, args.beta_fast, args.beta_slow) - self.register_buffer("freqs_cis", freqs_cis, persistent=False) - - def forward(self, x: torch.Tensor, start_pos: int): - bsz, seqlen, _ = x.size() - freqs_cis = self.freqs_cis[start_pos:start_pos+seqlen] - win = self.window_size - ratio = self.compress_ratio - rd = self.rope_head_dim - if self.compress_ratio and self.compressor.kv_cache is None: - self.compressor.kv_cache = self.kv_cache[:, win:] - self.compressor.freqs_cis = self.freqs_cis - if self.indexer is not None: - self.indexer.freqs_cis = self.freqs_cis - # q - qr = q = self.q_norm(self.wq_a(x)) - q = self.wq_b(q).unflatten(-1, (self.n_local_heads, self.head_dim)) - q *= torch.rsqrt(q.square().mean(-1, keepdim=True) + self.eps) - apply_rotary_emb(q[..., -rd:], freqs_cis) - - # win kv & topk_idxs - kv = self.wkv(x) - kv = self.kv_norm(kv) - apply_rotary_emb(kv[..., -rd:], freqs_cis) - # FP8-simulate non-rope dims to match QAT; rope dims stay bf16 for positional precision - act_quant(kv[..., :-rd], 64, scale_fmt, scale_dtype, True) - topk_idxs = get_window_topk_idxs(win, bsz, seqlen, start_pos) - if self.compress_ratio: - offset = kv.size(1) if start_pos == 0 else win - if self.indexer is not None: - compress_topk_idxs = self.indexer(x, qr, start_pos, offset) - else: - compress_topk_idxs = get_compress_topk_idxs(ratio, bsz, seqlen, start_pos, offset) - topk_idxs = torch.cat([topk_idxs, compress_topk_idxs], dim=-1) - topk_idxs = topk_idxs.int() - - # compress kv & attn - if start_pos == 0: - if seqlen <= win: - self.kv_cache[:bsz, :seqlen] = kv - else: - cutoff = seqlen % win - self.kv_cache[:bsz, cutoff: win], self.kv_cache[:bsz, :cutoff] = kv[:, -win:].split([win - cutoff, cutoff], dim=1) - if self.compress_ratio: - if (kv_compress := self.compressor(x, start_pos)) is not None: - kv = torch.cat([kv, kv_compress], dim=1) - # We performed QAT here, kv could also use fp8 format, though current implementation uses bf16 - o = sparse_attn(q, kv, self.attn_sink, topk_idxs, self.softmax_scale) - else: - self.kv_cache[:bsz, start_pos % win] = kv.squeeze(1) - if self.compress_ratio: - self.compressor(x, start_pos) - o = sparse_attn(q, self.kv_cache[:bsz], self.attn_sink, topk_idxs, self.softmax_scale) - apply_rotary_emb(o[..., -rd:], freqs_cis, True) - - # o - o = o.view(bsz, seqlen, self.n_local_groups, -1) - wo_a = self.wo_a.weight.view(self.n_local_groups, self.o_lora_rank, -1) - # NOTE: wo_a is FP8 in checkpoint; could do FP8 einsum here for better perf, - # but using BF16 for simplicity. - o = torch.einsum("bsgd,grd->bsgr", o, wo_a) - x = self.wo_b(o.flatten(2)) - return x - - -class Gate(nn.Module): - """MoE gating: computes expert routing scores and selects top-k experts. - Supports hash-based routing (first n_hash_layers) where expert indices are - predetermined per token ID, and score-based routing (remaining layers).""" - def __init__(self, layer_id: int, args: ModelArgs): - super().__init__() - self.dim = args.dim - self.topk = args.n_activated_experts - self.score_func = args.score_func - self.route_scale = args.route_scale - self.hash = layer_id < args.n_hash_layers - self.weight = nn.Parameter(torch.empty(args.n_routed_experts, args.dim)) - if self.hash: - self.tid2eid = nn.Parameter(torch.empty(args.vocab_size, args.n_activated_experts, dtype=torch.int32), requires_grad=False) - self.bias = None - else: - self.bias = nn.Parameter(torch.empty(args.n_routed_experts, dtype=torch.float32)) - - def forward(self, x: torch.Tensor, input_ids: Optional[torch.Tensor] = None) -> Tuple[torch.Tensor, torch.Tensor]: - scores = linear(x.float(), self.weight.float()) - if self.score_func == "softmax": - scores = scores.softmax(dim=-1) - elif self.score_func == "sigmoid": - scores = scores.sigmoid() - else: - scores = F.softplus(scores).sqrt() - original_scores = scores - # Bias shifts scores for expert selection (topk) but does not affect routing weights. - if self.bias is not None: - scores = scores + self.bias - if self.hash: - indices = self.tid2eid[input_ids] - else: - indices = scores.topk(self.topk, dim=-1)[1] - weights = original_scores.gather(1, indices) - if self.score_func != "softmax": - weights /= weights.sum(dim=-1, keepdim=True) - weights *= self.route_scale - return weights, indices - - -class Expert(nn.Module): - """Single MoE expert: SwiGLU FFN (w1, w2, w3). Computation in float32 for stability.""" - def __init__(self, dim: int, inter_dim: int, dtype=None, swiglu_limit=0): - super().__init__() - self.w1 = Linear(dim, inter_dim, dtype=dtype) - self.w2 = Linear(inter_dim, dim, dtype=dtype) - self.w3 = Linear(dim, inter_dim, dtype=dtype) - self.swiglu_limit = swiglu_limit - - def forward(self, x: torch.Tensor, weights: Optional[torch.Tensor] = None) -> torch.Tensor: - dtype = x.dtype - gate = self.w1(x).float() - up = self.w3(x).float() - if self.swiglu_limit > 0: - up = torch.clamp(up, min=-self.swiglu_limit, max=self.swiglu_limit) - gate = torch.clamp(gate, max=self.swiglu_limit) - x = F.silu(gate) * up - if weights is not None: - x = weights * x - return self.w2(x.to(dtype)) - - -class MoE(nn.Module): - """Mixture-of-Experts: gate routes each token to top-k routed experts + 1 shared expert. - Experts are sharded across TP ranks; each rank handles n_routed_experts // world_size experts.""" - def __init__(self, layer_id: int, args: ModelArgs): - super().__init__() - self.layer_id = layer_id - self.dim = args.dim - assert args.n_routed_experts % world_size == 0, f"Number of experts must be divisible by world size (world_size={world_size})" - self.n_routed_experts = args.n_routed_experts - self.n_local_experts = args.n_routed_experts // world_size - self.n_activated_experts = args.n_activated_experts - self.experts_start_idx = rank * self.n_local_experts - self.experts_end_idx = self.experts_start_idx + self.n_local_experts - self.gate = Gate(layer_id, args) - expert_dtype = torch.float4_e2m1fn_x2 if args.expert_dtype == "fp4" else None - self.experts = nn.ModuleList([Expert(args.dim, args.moe_inter_dim, dtype=expert_dtype, swiglu_limit=args.swiglu_limit) if self.experts_start_idx <= i < self.experts_end_idx else None - for i in range(self.n_routed_experts)]) - assert args.n_shared_experts == 1 - self.shared_experts = Expert(args.dim, args.moe_inter_dim, swiglu_limit=args.swiglu_limit) - - def forward(self, x: torch.Tensor, input_ids: torch.Tensor) -> torch.Tensor: - shape = x.size() - x = x.view(-1, self.dim) - weights, indices = self.gate(x, input_ids.flatten()) - y = torch.zeros_like(x, dtype=torch.float32) - counts = torch.bincount(indices.flatten(), minlength=self.n_routed_experts).tolist() - for i in range(self.experts_start_idx, self.experts_end_idx): - if counts[i] == 0: - continue - expert = self.experts[i] - idx, top = torch.where(indices == i) - y[idx] += expert(x[idx], weights[idx, top, None]) - if world_size > 1: - dist.all_reduce(y) - y += self.shared_experts(x) - return y.type_as(x).view(shape) - - -class Block(nn.Module): - """Transformer block with Hyper-Connections (HC) mixing. - Instead of a simple residual, HC maintains `hc_mult` copies of the hidden state. - hc_pre: reduces hc copies -> 1 via learned weighted sum (pre-weights from Sinkhorn). - hc_post: expands 1 -> hc copies via learned post-weights + combination matrix.""" - def __init__(self, layer_id: int, args: ModelArgs): - super().__init__() - self.layer_id = layer_id - self.norm_eps = args.norm_eps - self.attn = Attention(layer_id, args) - self.ffn = MoE(layer_id, args) - self.attn_norm = RMSNorm(args.dim, self.norm_eps) - self.ffn_norm = RMSNorm(args.dim, self.norm_eps) - self.hc_mult = hc_mult = args.hc_mult - self.hc_sinkhorn_iters = args.hc_sinkhorn_iters - self.hc_eps = args.hc_eps - mix_hc = (2 + hc_mult) * hc_mult - hc_dim = hc_mult * args.dim - with set_dtype(torch.float32): - self.hc_attn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim)) - self.hc_ffn_fn = nn.Parameter(torch.empty(mix_hc, hc_dim)) - self.hc_attn_base = nn.Parameter(torch.empty(mix_hc)) - self.hc_ffn_base = nn.Parameter(torch.empty(mix_hc)) - self.hc_attn_scale = nn.Parameter(torch.empty(3)) - self.hc_ffn_scale = nn.Parameter(torch.empty(3)) - - def hc_pre(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor): - # x: [b,s,hc,d], hc_fn: [mix_hc,hc*d], hc_scale: [3], hc_base: [mix_hc], y: [b,s,hc,d] - shape, dtype = x.size(), x.dtype - x = x.flatten(2).float() - rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps) - mixes = F.linear(x, hc_fn) * rsqrt - pre, post, comb = hc_split_sinkhorn(mixes, hc_scale, hc_base, self.hc_mult, self.hc_sinkhorn_iters, self.hc_eps) - y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=2) - return y.to(dtype), post, comb - - def hc_post(self, x: torch.Tensor, residual: torch.Tensor, post: torch.Tensor, comb: torch.Tensor): - # x: [b,s,d], residual: [b,s,hc,d], post: [b,s,hc], comb: [b,s,hc,hc], y: [b,s,hc,d] - y = post.unsqueeze(-1) * x.unsqueeze(-2) + torch.sum(comb.unsqueeze(-1) * residual.unsqueeze(-2), dim=2) - return y.type_as(x) - - def forward(self, x: torch.Tensor, start_pos: int, input_ids: Optional[torch.Tensor]) -> torch.Tensor: - residual = x - x, post, comb = self.hc_pre(x, self.hc_attn_fn, self.hc_attn_scale, self.hc_attn_base) - x = self.attn_norm(x) - x = self.attn(x, start_pos) - x = self.hc_post(x, residual, post, comb) - - residual = x - x, post, comb = self.hc_pre(x, self.hc_ffn_fn, self.hc_ffn_scale, self.hc_ffn_base) - x = self.ffn_norm(x) - x = self.ffn(x, input_ids) - x = self.hc_post(x, residual, post, comb) - return x - - -class ParallelHead(nn.Module): - - def __init__(self, vocab_size: int, dim: int, norm_eps: float = 1e-6, hc_eps: float = 1e-6): - super().__init__() - self.vocab_size = vocab_size - self.dim = dim - self.norm_eps = norm_eps - self.hc_eps = hc_eps - self.part_vocab_size = (vocab_size // world_size) - # lm_head in the checkpoint is stored in bf16, while the parameter here is stored in fp32 for easier computation of logits later. - self.weight = nn.Parameter(torch.empty(self.part_vocab_size, self.dim, dtype=torch.float32)) - - def get_logits(self, x): - return F.linear(x[:, -1].float(), self.weight) - - def forward(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor, norm: RMSNorm): - # x: [b,s,hc,d] - x = self.hc_head(x, hc_fn, hc_scale, hc_base) - logits = self.get_logits(norm(x)) - if world_size > 1: - all_logits = [torch.empty_like(logits) for _ in range(world_size)] - dist.all_gather(all_logits, logits) - logits = torch.cat(all_logits, dim=-1) - return logits - - def hc_head(self, x: torch.Tensor, hc_fn: torch.Tensor, hc_scale: torch.Tensor, hc_base: torch.Tensor): - shape, dtype = x.size(), x.dtype - x = x.flatten(2).float() - rsqrt = torch.rsqrt(x.square().mean(-1, keepdim=True) + self.norm_eps) - mixes = F.linear(x, hc_fn) * rsqrt - pre = torch.sigmoid(mixes * hc_scale + hc_base) + self.hc_eps - y = torch.sum(pre.unsqueeze(-1) * x.view(shape), dim=2) - return y.to(dtype) - - -class MTPBlock(Block): - - def __init__(self, layer_id: int, args: ModelArgs): - super().__init__(layer_id, args) - self.e_proj = Linear(args.dim, args.dim) - self.h_proj = Linear(args.dim, args.dim) - self.enorm = RMSNorm(args.dim, args.norm_eps) - self.hnorm = RMSNorm(args.dim, args.norm_eps) - self.norm = RMSNorm(args.dim, args.norm_eps) - self.hc_mult = hc_mult = args.hc_mult - hc_dim = hc_mult * args.dim - with set_dtype(torch.float32): - self.hc_head_fn = nn.Parameter(torch.empty(hc_mult, hc_dim)) - self.hc_head_base = nn.Parameter(torch.empty(hc_mult)) - self.hc_head_scale = nn.Parameter(torch.empty(1)) - self.embed: ParallelEmbedding = None - self.head: ParallelHead = None - - @torch.inference_mode() - def forward(self, x: torch.Tensor, start_pos: int, input_ids: torch.Tensor) -> torch.Tensor: - # x: [b,s,hc,d] - assert self.embed is not None and self.head is not None - e = self.embed(input_ids) - e = self.enorm(e) - x = self.hnorm(x) - x = self.e_proj(e).unsqueeze(2) + self.h_proj(x) - x = super().forward(x, start_pos, input_ids) - logits = self.head(x, self.hc_head_fn, self.hc_head_scale, self.hc_head_base, self.norm) - return logits - - -class Transformer(nn.Module): - """Full DeepSeek-V4 model: embed -> HC-expand -> N blocks -> HC-head -> logits. - Sets global state (world_size, rank, default_dtype, scale_fmt, scale_dtype) in __init__.""" - def __init__(self, args: ModelArgs): - global world_size, rank, default_dtype, scale_fmt, scale_dtype - world_size = dist.get_world_size() if dist.is_initialized() else 1 - rank = dist.get_rank() if dist.is_initialized() else 0 - default_dtype = torch.float8_e4m3fn if args.dtype == "fp8" else torch.bfloat16 - scale_fmt = "ue8m0" if args.scale_dtype == "fp8" else args.scale_fmt - scale_dtype = torch.float8_e8m0fnu if args.scale_dtype == "fp8" else torch.float32 - super().__init__() - self.max_seq_len = args.max_seq_len - self.norm_eps = args.norm_eps - self.hc_eps = args.hc_eps - self.embed = ParallelEmbedding(args.vocab_size, args.dim) - self.layers = torch.nn.ModuleList() - for layer_id in range(args.n_layers): - self.layers.append(Block(layer_id, args)) - self.norm = RMSNorm(args.dim, self.norm_eps) - self.head = ParallelHead(args.vocab_size, args.dim, self.norm_eps, self.hc_eps) - self.mtp = torch.nn.ModuleList() - for layer_id in range(args.n_mtp_layers): - self.mtp.append(MTPBlock(args.n_layers + layer_id, args)) - self.mtp[-1].embed = self.embed - self.mtp[-1].head = self.head - self.hc_mult = hc_mult = args.hc_mult - hc_dim = hc_mult * args.dim - with set_dtype(torch.float32): - self.hc_head_fn = nn.Parameter(torch.empty(hc_mult, hc_dim)) - self.hc_head_base = nn.Parameter(torch.empty(hc_mult)) - self.hc_head_scale = nn.Parameter(torch.empty(1)) - - @torch.inference_mode() - def forward(self, input_ids: torch.Tensor, start_pos: int = 0): - h = self.embed(input_ids) - # Expand to hc_mult copies for Hyper-Connections - h = h.unsqueeze(2).repeat(1, 1, self.hc_mult, 1) - for layer in self.layers: - h = layer(h, start_pos, input_ids) - logits = self.head(h, self.hc_head_fn, self.hc_head_scale, self.hc_head_base, self.norm) - return logits - - -if __name__ == "__main__": - torch.set_default_dtype(torch.bfloat16) - torch.set_default_device("cuda") - torch.manual_seed(0) - args = ModelArgs(n_hash_layers=0) - x = torch.randint(0, args.vocab_size, (2, 128)) - model = Transformer(args) - - print(model(x).size()) - for i in range(128, 150): - print(i, model(x[:, 0:1], i).size()) - - h = torch.randn(2, 128, args.hc_mult, args.dim) - mtp = model.mtp[0] - print(mtp(h, 0, x).size()) - print(mtp(h[:, 0:1], 1, x[:, 0:1]).size()) diff --git a/TEMP/dsv4thing/requirements.txt b/TEMP/dsv4thing/requirements.txt deleted file mode 100644 index 7e1cc78c..00000000 --- a/TEMP/dsv4thing/requirements.txt +++ /dev/null @@ -1,5 +0,0 @@ -torch>=2.10.0 -transformers>=5.0.0 -safetensors>=0.7.0 -fast_hadamard_transform -tilelang==0.1.8 \ No newline at end of file diff --git a/CORRECTNESS_BACKLOG.md b/archived_plans/CORRECTNESS_BACKLOG.md similarity index 100% rename from CORRECTNESS_BACKLOG.md rename to archived_plans/CORRECTNESS_BACKLOG.md diff --git a/DEGENERATION_TESTS.md b/archived_plans/DEGENERATION_TESTS.md similarity index 100% rename from DEGENERATION_TESTS.md rename to archived_plans/DEGENERATION_TESTS.md diff --git a/reference/official_inference/README.md b/reference/official_inference/README.md new file mode 100644 index 00000000..ad17929c --- /dev/null +++ b/reference/official_inference/README.md @@ -0,0 +1 @@ +# THIS WAS FROM https://huggingface.co/deepseek-ai/DeepSeek-V4-Pro/tree/main/inference IT WAS USED TO REFERENCE HOW THE PARSERS, TOKENIZERS, AND TEMPLATING ARE HOOKED UP. IGNORE THE KERNEL AS OUR VERSION OF DSV4 IS OUR OWN NVFP4 QUANT \ No newline at end of file