From 4523d0ce373ee4b2176b3251fff29fd4864fcf38 Mon Sep 17 00:00:00 2001 From: Bhargav Krish Date: Wed, 29 Jul 2026 21:59:23 -0700 Subject: [PATCH] parakeet : verify hparams loaded from parakeet model bin file (#3950) * verify hparams loaded from parakeet model bin file * flexible way to accommodate CI as well security concern & test case addition. * add bad model for CI tests * removing whitespaces,couple of nits --- .../for-tests-ggml-parakeet-tdt-bad-nfft0.bin | Bin 0 -> 14299 bytes models/generate-parakeet-test-model.py | 13 +++- src/parakeet-arch.h | 51 +++++++++++++ src/parakeet.cpp | 68 ++++++++++++++---- tests/CMakeLists.txt | 1 + tests/test-parakeet.cpp | 27 ++++++- 6 files changed, 142 insertions(+), 18 deletions(-) create mode 100644 models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin diff --git a/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin b/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin new file mode 100644 index 0000000000000000000000000000000000000000..bba7d72b22134f9db55bc5b82afcc792c4c460c5 GIT binary patch literal 14299 zcmbt*2|Sfu*SDbz4TfY$2$fJ2MTT?jby8{2+#Q9I42dWtX+o$B8A521kTOJxGF*Ed zO-K_BhExhkq9oCHe~zc#x95BM-r@Q7@3;23&a}?j|F!miuf48)HoI@!I7~o5AYg=m zfY`vJ(0~XIh}3}mYW&xS42bZ6@UQvr^FKFqK=_~IKYt;7oJ>~}CokQt zUQTxZxYpg`vrX66Y12lRZEd$p$tph|RC&1`l>Ha6QHvPryj&MKyN|;cjtZ?@{Djso zyhL5+iox!&B_wb819CX7k(5b^b5msRQ!900+_Z8TKFm`F?S0X(U~N0=KiUeXORq%k zAq{vvmw@Jo;ZV0u8}m!tNq>JA)Ar7e^vc?S=_f}r?~V?%I)^bnqH$#GaW#+)F2)a8 zsWkWDdm3H4C# zbZMFm%x|4d^A@CHq+|@5*t=mv_zyC&PMVy#HVqni4b*hCI=P^s3RxRYQa);92ZZ1L z3$Qk9viJ6S)=s{7YoY$jBC6qijhag*q0e8RN#LsYRQ=K-4&P0xd_4n=?>WBk)$?Ef zgAW|Gcx>05_R|gS6F!mB*(Uysch_*TKsBKgvLwGVT^@88sy$4 zJMV17{U^()^5QdK>QYE+?r6h{ zvqBN3UWsKMM~Gs!nFpCGoX?RLy#r}ud}yCw5Q@dkB=^njN#o}n;M=S+ApFK3!>@7h zAOGKD;iJ@0xeD%HE+qRFmwJ6J>do;eJHs6(#g|j=%siZxqz-E$9jN{6 zQuO~3h|R&L$o42pcoFwE6zw+Q`sq9*b}ROvm$w*ir*tTJt+Wga&OM~-FDl{g$x?Kf z*FmtVRwYwB)#>8sU=+4}jMm@#se*+qd%)r)#u`>N9D!ERo^w=j2 zcJir780*kN5_?|LGfl_QCpHzEi>HyyakrrM;dqkK|AD=0b(ly`&Vak=>xqO@4Sja> z(7)x|>4T%_3wc5^Vz0x!6jQRXO2LXEChyWFVxgkO_2RrA#8-9Dv;TA78vaB@)599u$H!9X6(tyJ zB#&Q@q>xQ(jslZzf)ln+!l~}!yw(^Es{j1}t_~H&9>4WiJfe&x-6`l$GEeVD742O#@o)1%&6@v9&HOoZF6kLLkietOV&e43 zw_L1l_NLtq2k8gaopfqRAb;)C`Ax6t!1XV#`VXuB+2GBaJe=&kei`oX-c7Pk1;R7U zam01ZC9=6_4Ws`p&q(x&5rz&wP4r1QyW`VP?nQ@6IDhB{xneaN0~{jJ&NmUH59}v4 zk9tYslIM7QvpZCrdrcD6BQf8tn|2R9K-}9bne7Sj^lbM7R#!NXt}#|3^=nE&ea-}4 z(?KtqjWfvphy;@0MyXk26NCw!Dq%<%1qtG+TYlZwM-Um+;|3B zHD$q}^eJp9I0F6mng|#_hr(-3M67QmXw)UpDGLrV>%#QWZq**hs5yel^XG$@$!zN0 z5f42^TzEIE9Y&5eg%d`Vh$^A%sv$@4d(IGix1|hR4aURRPJP^47Kp*-A`H254D0Vr z#sW+xb6&(kRqa^VlQxRi`{)?4+8@ifw?81g&+o#R5kq*9E95bCt~}&=-ojTMrNAVW zlH#J(uvTjnRTC9;)xxP}q^A#3rfv^;&Lt3;%dLnSo`w$6wbO}rAWOz?g)-fp& z-Sota4#LTw4q8zoiEN@7I25{)6&FQehtmv{csrA3`!67qBvycDVkI?l8ww&y2Wd~y zd)C(GFF307H{Sh_j@&CO_N|fSv8Q$LRRh4v#u&OgbRpU077oeX=Rhy<5WCnc2imer znG!ooOx9LFqkSw=TLqK0@{S`amz{5FEAGj&&jzNz~^1bbMkA6}rg~fqg2lFh7Sbm~s%r zCLhN3%lAwda5cC?wm6bl@5|8TQNrBI84dpNF6>HwU8)r^54SICBvwDxV3mV5T3>34~7`f26%X5y-E6<%-P?6yZAH*DslH{@rZqmX{C1;$G5sq!^6SL`aNbFgZWJ z0=^m-QTxuVXpq@TmnXK7D%T3M&>IgOZ^dzc$6<&z9m891y8x#>tHql^D^czE9y0a7 zOnRv}klG54r$Td=kwW=E9P@Do>`c>w3*&E)FFPba@p=MDBSKt{br)G-lM{4m8yBal z1!Kd)WHdhLfK9w;wnQ_FSngcTjQcVh{k@z)V<7|eJIB+YiIecA(+JSm%_SV$BIFo{ zvr$_6iSr6)T;gm-PR%UFiDFsspfnjiE(oM~a>Yc8qsW^s9EWImDoelHdTS-P)EuJbB2Ls1oG=964ba4vNGPs4z z*|P-gbkeY+)DX^j`jU7ZWi;g#f&PgX)G|nz5y=fD?x$6-+^CzZiWWn&t2I={?L55I z3&9_*ioB`Q@4__A@!W0x_vq3$eOT6al2(TwrcS%eXh2XTs;nCgD&``PJ0CIXhd$`p z33Kh_rjeHg*>vG1X}Tj&8P4o~isF`KGX(MqD;8m+q0zVDGPJz?M6A(YN{&X}LcXZ%mlL z3+tE=`Dl_oo{IOu6kN>_Y`S6yNli`@&*-UDPH&KmAG}$e7qbmnkRH& zJ7&*wMfhAncL`YAuz^OEO@yi=cBQ7L zri166V9ft9w^S9a=(5wtaD-~bz$S1ih~?!|*@k_%{quG9?S^PL-Y3m_5`(zFY%{it zmH^3(CT$*bpz3`Leq0oVwewnu^aT$_W9JBXFojE>xI~ivb~SQs>^x9NTmpSo_YfXP z^7M=YV7+J}+)Ukva*|zmqvjUlYT*HK@~>d@lp&C@oKk}{N1T4rlw)k1M6ojv@2d!N zoda56xN8hFl)7QWOKXUvuxC)}(sFjzf>d_5+$+%Ta2!mP#{EOzf2WY2A!US^aul!bCBqa9nB0n` z^m9ZQuFg@Uiv$3zhMgMBkf!{mj$gp|mkeppyscX{d2I9D~BSIM06G*G`C0X@#5bi7$I zcH5`YC!Xs_d)E^B?7a-C>w6&WiZs)+qY2i^AS*IMgzOksLae+jxDB0aV4@_2ghzAMQ)jhEq-1_esUS^`e-KS3_}>Hy3;6Rgn}E3D{sh4&!g`hM8J%P+E5xQV)p2qu}!(NZ${} z&*Y!+^Q$2C=lF4S+Pclm$% zwtD$^cx>|6Xy@)^??HR-!sp2%CK>fV7n?7`OuL^@BJeN-HV#RI*=b8jrn~?S+t^QzPZY(J(86L5@7}Bm5%3o zXQ2JjA^)tDALtFK+whoR6RGPyi8g)9s6bQ`BfLk6`@$j*Kf9I=Ce^>}#oI~m1WdcN zogHNuiqHCAljhG?pkwn%*5&vDdLt`xkmY`_eEboJ|FrlYH3U1)pKNlm%9Z9eRiWL* zO?cGAg}6UD3re@9b3L`?@Xh7Rz&+|mGMg0X&^=GdgyVadN%Ak5;nLOAtoI@Y#os_l z(@AhT>kyrvq<~XxhI7sbSis5)Mr3BsB%IuL9C+%t;Gp4mB5dXd5&b7YAtM@&87PCf z&KTG!HWZ?LA~12obXXRigU?+p(CEn-$WpKZX$K0g6~bxD?pT~`u$@-+-e(?PPhmWj zVqk0k75bF9M6%@e;J2t(r22}5>GdAKS$jNbOd^NgTpmW3mb?VXla*{6`q5O52=C?_ zM~u3#pAL(gid9eYz`IqG(vGx2xBTzXbNh*&0Nu^Z+^xX0Hx#g;w|1EB6l{dRZKHWB z#|ZGq$l+)$;=#$yoCdp%K9CnB*QvRhzi>7qa)t)aQXDzIBvZyZZ%it?n-_IHTP?1!lYuV zGA0&JG9U59x_7LJw?C@f%wZ!Q+{UX?v+Ei;z0h zE!Hjd)yc!mmK546RYDebD1o}26kR2K9h5%>z+{&o`c0Gz5&Qcw zO`*)^z`q%L+kc`by5Jmc?dqq!lWNJ@f@;<>Tmz>iO~E1?Lkyo}kLDtyxOOWONt%!Y ztabT9-Q5WF%QD1-0xjP04S8sNwGDl~N%Hm!CZR^+Imj0>WcBsZVd>@?a&y&un402@ zXVgSs$RQKZOqRi6L=w(-OTn((5!h51g~2fiSgmz|dF-A+KEGUu0hRTf%Mzmo0!amg z(r$v%wp3WI9e`#2wJ^F%6dh>@9Oi^GWatBWd_x-^J3o~^Pb$Z^r+i_y##OKtT7>P# zjuNwcV_e^&iVNnN!=eUdwlez}=(K&J`yTp3Za_Jsbj1SdyIbi0XifQ-nfsqN!wh~b zWOKsavbCa?sIylKu58zU={@qWu;e*M_%i+`0c6cUJ^IvNv^2n2fQvs?z@c&-l2g@0w++2T z%jbykyxN0Eyyt3Y&8ftbE41*=q&PHF^Tk7wHcXxKP~JnXGNgawLP&G~?YgGIvyk=2 zl6xoM+$}v`Le~X4XIKKR=wA$TA5KB32Y1Pkt5V=C-T}h*m(u2gJhDiRqMuzN6y%xV zkwcOwG^?;QaYO}PFmol(vwUIm*Ld>x`l`X``)}9ef5C5nvVX>3anyJAST6(b?cd|O zm{8PmS`R%7!pQs18qz7;1sbp3ku`3b@MO0RH~yV99d5P=)x@5VEdxc%n#Uq|WMVMx zeJH~fU7bNBd`)rFPFL7_;5nP)7f6-6CBfM?9V5dpkxNfLLhbo3^86%^W>gEJtU?WS z65U4|x0F$bSXa1RKd=|iwgq$Nt90?-JThTcG{#G1Lf_V8D9-SK-OXh>bb(6{Xp@-r>X0pYhlu19~+%wI+)uT6keFVbn!z{V7^YvxDgrAnpjg<-Y#W7?T7!ToTd0-c;wAVEGCMo+Rv zw`>XSwz|C#b9@I>Mz(T7x2~gewph|rJ>_JSbv|s#G$wX`ZNiL^D=Cn4bB{ez0JRiuL2uW5j2hgW2QffPX28^*TlEOC>b2mU5Y z(W_!TMAKZXA9t2K7oABq+)&0HkBdsH){Vn`I-&SuL;($0dI$^WrV|0fSX7;32n8ky zc&T3sg@2{qstN$iUqX@4)dm0;l{%@!Y5hapx;mm{;Bf!Dj3BG6l zjeeH3c+}F0Xb856U#9Yo(Jk~w8<*;=tfvJEVK8P>5-i!N zOKVPw@w7!$Kz@n@9CkTL%*=Ly+?aUujuPNXp724LI{-poYw_Vk4_v2k70$(l(8)us zU`v`8S^V)ad*<+YxM1l;P2M%*#FWL@rZ^HN2iziS<0~-8!x%4gjK)VQCG@>=9(FaV z;>`UfI3s8-#Hud^%cZi^v#5fum@i4Bhc8E`s>5h`aVx3>sbk^WIPCr~k)+=>fngd| zSln=jQ*UTbO&ZP<)$c1y!zGLHb9@RGig$Aqa)SU0M)G8YFObl=iBNq}iH>R*0hZsU zK zkiK=+Wo`RUF$za?FvVC5zU?#NF51?N{VnOF`j{D>y)OeJ^9l1tY&Hn^9)aP*<`Vg~ zX!s%~$Lo%*f{9{0a^q7dG*38=mN|c6{Z?5rJT|;^T~{=5styD8!wWQ6K zWi)VcDc&v@02-jk)A&9SR(_|L`^lQAlK9BF=f{J6w;`3z*Tc(sg7~2GG+bRM4i=s_ zsPH&1YL--i@~>`^v{^5()Q1c91;V8IngW+pM_?eghPBXr^-l!CCyzgp)1L_B7xVe} zC&EV)eY|4jL{5ten}*m%lcGj_JlR>uevP<|ch(LgN}FDh(@ie)d(mBVm)ZlOsl`OK z(HNg*M$vb&JWS&Tq2PIOY~`E)|HAP&ymdUt-xJ5a);;7~r4&sKumh)dF+4WE5(}F1 zP*qC|^(hmYSn@p$s->pCZd| zex=h7JA#F|InMW1!MCECn14qKmtKs3s(w}QdfouCYSP@$t{`li8O@X&I)YQB=b_wU zHEdXR07=?@T!1pvgFS|)+)f%66f$^fb5$$6(&)Lm*pTg;GvO$;NMiU^V9< zb_WWUjxvZK)!&n7noK9Qm7Sq+aT(O&lswG@Dcn=@l1ikivztO+k_S(Z(LlDEDo5{j{-cnf^ZE!2|fCIPt z$kOl-x}imcm-%)M9J4hwG;Y!Eziiz8{7gHDtdYcfR2lo55;2{75F5qb zpyXo{2nuThUC%-Y9e*EXqKzS5v=IJ&*-o>s-D6e+3h>fn4w7Q^gE-E85%GUij#vAK zf$8`l*tvH;p5GsX^S_DnY&IUGR?VaF?tz;q^P!D|Yh$p$*>h_i9uE?Nljb7Qjkvs-7oBzzwWB(4vGoC@v*wc#@0|(C7i}viMQ;8@xPZnKf z>>Jeg>x|BSt&g7{z2mI~#{g6C8`+4bQt#5z&$d_^lMW;9XQREJ8y|#sx4Xcn?j(#q zd6@1JmBXCbnnbg2JRZInM&)jM(#HavfumpjL{@1VG1*nXiaRw9+{ChggL^%xyXXm8 zyzwzprd1DB(OGCEI~TUB{D4n(uOrKjSHd;@FU0DKE;;EZiTRd4IQ9~O^qRo}jB|~n z9uq8xrOO=L?Xm=Y1~MO|kriag(c$1^Vc$24V$P+J zTovaw`f0m1h8jL2Qau60@P0nN&5r`Pb{_j(KN3f4CxZDDRS0TK!Q!3G@ZpRm93MHU zbYVml^s0;}J0)I~dWdVG#>dgHWavlIAH0f>u>MG6NLAHJo=D*o_o!(ME9c!a3vZw8}&O^{yF3s<+!#!mPXF5e4gl`7~2EmtK>-gfBI|qi075mD#kBX-`hW zV1tqDnaQ?{_`sw+KGzjngYT2*DK|-SRxBED)d}-#AI$192mOiH*o33+=oKY3l-d7| zSXRy>$*OA5B`!*Wk1mIjqB6$z#e5)L>%gaXG0yR9L!Fo#aMLA-m~VT6O{InynKT#f z=tWTnln3z-fTF zC3@($(i4TQZGm%9#dN8g5P2-MkCmHk3i3CNIJX-rSOr53Y}KGxbeCm%FU~=0XL&IjXF zg&cPN<0%6N4;6T0R|Z3vbpkH?dXDTnx|c&9MtW08Sv=q;5(Sa;?f{{T(@% zNbmDRxrGv2$t$NWtX~{}cPgHtwO}M?h|gzs@1u8AIA4G}A#4Q1zc7YJ)*kq&+#lDL z$uX6hdgSD>PTY9fh$colbEKBpQX7+H_(f$H#wi-mwY^15@&iTiUlRs+=?GnNsGYiu zZX&DJD?!N8XVk=MA&kAb6TYrWhoPsM!TZiyTwdV7`TyV#7e0R+Q922ms;xN}JPyHu zJBP^ZiZ^6UV-ZP=)Q8Pw(%k2GC9wy-(Zj`kV4`8e@#LmMYh4KOh*~*N zfQlrRGjidasuL3`{Db_6Dj}1O#?rbq_sCQMN$w8MFsgX{F=_j83Tk>{z~~cjJVODM7n?@-_h3SygH~vue-Wg-A+{C275OOV5oMp1U($)<*OqQ}9*Ee+)^sm|v z+!3o`h#d?|#D4$w1-BH>00ZBL_+IfOc{t4h zPx!>4nAuBebJ3JOdNTpd3gp1Rp@rV++zDO1QA8SX^z4CCTH!EZ*TDKI`h`8cVgzsBk`X9({sed) z83_i>dk11E5ati-CvMMz@ zVc?F#Q2NVZSbAp|ygIa#JQ9>8HQLi@Xl)3LtdZv(pCFEtGoA5nZY7-gSdUZpUcucm zk>Hw{12ZlYTz_8;kCrLJjCZ4XiKbn+ur!L^-*6L#C_JYTI|t4_TJTscn3FpE zEzQ3Fl=4e)S_8swssr+ebnu^YkHNcHeBubxpUmSNh{e1rBW_*6X;>}&1srC`1ADZR zmL|o+b$c!^T#?qta1zrc678GN&gqH?2cq3MP^ z_sx+>D6RJdd73QbJ*wwe9t+3Mq1{CBNf7kPw<5>3WMJYeLT1P&E}!_M2ZZ1FnK55n z`R6?N{|xz+y;>U#NAWY4K(9HUO*y%od}})bW3D=(!Q{R8(r^JyDq9KDtTRz_-y>ZA zd<9h6Dq?+g1?|h~#Dr=cl*G2V;>mVGun(SRZa^RMf z3JuoFfza#{upP4;W$x*)6Vt9RraO0|u$&|7=yu+8!-p7Xi#Q7Ek!R^aIU6b|sSWLS zWpUgEIan@si4D}=f!v;LXk=UjAE&6&aedX0xkni572?TBwJ}`p)d%oM#|Qj&J`Jv& zTZvAm&Xtz_kj2xBLWz~7H_kLOqWVi-kvURBUxs6`KRKIS%U><*(3#ZpV}KAVI)+oys0z$D{i{Oj*S|BUfKRrTivVa~Nq+UC|z zrmyp%nrbD)*R2AAW(vW;77bOrUXxvSPr%qW8|d2?dtp=7G}IVKQjSMT5z~iFcyPZD zh0AH6AOTw5S`mHE8FX{+ZK5)MCNH7vD4C$~nsv8Z$}WFC1H4s}*hL>=@M)D4J!l=q z?AXx3=-r)2h7St?UEy%7c{823DWwBX^cgPhok`X`kwxAB$N8pg#7if2=$!K%@T4vU zD+=PVD*X%&S9%J!F2AKh;x(l)#&-uNki;OszeCK&hmUan4NqKEwUs_ysm!&XdJ)vE z2a3q$6Cvl>UpQoaGx8Bu`0d*x2d@7S#e=<_htoFQU$v!G{`=tVo;VoRI$-(X_tdOh zo2I@8nCw+S8p0|mGWXDVW;^p;DGXv4CShsFczA^8$YpUZR2tiXUxgkX=)a4labj41 z!U1mZhT(&3Jrq)J$1=+%=(k@9ZSr#m&XM_oR;moF*4s(TFHEAB+Z$+NvL=_iOai#R zoiuS(7n$T3h&i!5cwnl9LEr}2?N`Q_e4WOr(A~%eotD81>aVDg0}I>c zKVv^=Ex>_uZ1|xr0wO<7fl^ONsHqR4FUv>4X5Iu;dmBv0&fNj$?%RRax?NZyY=fNS zXUyyFMqu=M@oZounWmnH!4B%U_DC^(zNnTqoDstrT5nDFiHPG}y%G#fp9iL;Trh92 zf?I|+@Xqig(lbH_JbnbBqlFyDK)H)etm=T~Ri7Y5{1FBmAIIe^az&SsyCBkQ2Sm2+ z#iT>uNQOxxd-(N8ur=6BGSrel#_cGh+cud#vWS3(aid|}YE!UQ*TyMTA;?%tqehp? z;AE5g>B~R+kslE@a~6PngCFc{P@y)OmqAcF7h)Jc60mUvcBfauU!8uex1Jy=Y|5dF LDjjhBw-Wq6o%g1* literal 0 HcmV?d00001 diff --git a/models/generate-parakeet-test-model.py b/models/generate-parakeet-test-model.py index 192a96ce6..8b31a042f 100755 --- a/models/generate-parakeet-test-model.py +++ b/models/generate-parakeet-test-model.py @@ -3,6 +3,7 @@ import struct import sys import numpy as np from pathlib import Path +import argparse def write_tensor(fout, name, data): n_dims = len(data.shape) @@ -16,7 +17,7 @@ def write_tensor(fout, name, data): fout.write(name_bytes) data.tofile(fout) -def generate(output_path): +def generate(output_path, n_fft_override=None): rng = np.random.default_rng(42) hparams = { @@ -37,6 +38,9 @@ def generate(output_path): 'n_max_tokens': 5, } + if n_fft_override is not None: + hparams['n_fft'] = n_fft_override + n_vocab = hparams['n_vocab'] n_state = hparams['n_audio_state'] n_head = hparams['n_audio_head'] @@ -178,5 +182,8 @@ def generate(output_path): print(f"Generated {output_path} ({size / 1024:.1f} KB)") if __name__ == '__main__': - output = sys.argv[1] if len(sys.argv) > 1 else 'models/for-tests-ggml-parakeet-tdt.bin' - generate(output) + parser = argparse.ArgumentParser() + parser.add_argument('output', nargs='?', default='models/for-tests-ggml-parakeet-tdt.bin') + parser.add_argument('--n-fft',type=int, default=None) + args = parser.parse_args() + generate(args.output, args.n_fft) \ No newline at end of file diff --git a/src/parakeet-arch.h b/src/parakeet-arch.h index 3407a95c9..e8c6effe4 100644 --- a/src/parakeet-arch.h +++ b/src/parakeet-arch.h @@ -65,6 +65,23 @@ enum parakeet_tensor { PARAKEET_TENSOR_JOINT_NET_BIAS, }; +enum parakeet_hparam { + PARAKEET_HPARAM_N_VOCAB, + PARAKEET_HPARAM_N_AUDIO_CTX, + PARAKEET_HPARAM_N_AUDIO_STATE, + PARAKEET_HPARAM_N_AUDIO_HEAD, + PARAKEET_HPARAM_N_AUDIO_LAYER, + PARAKEET_HPARAM_N_MELS, + PARAKEET_HPARAM_N_FFT, + PARAKEET_HPARAM_SUBSAMPLING_FACTOR, + PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS, + PARAKEET_HPARAM_N_CONV_KERNEL, + PARAKEET_HPARAM_N_PRED_DIM, + PARAKEET_HPARAM_N_PRED_LAYERS, + PARAKEET_HPARAM_N_TDT_DURATIONS, + PARAKEET_HPARAM_N_MAX_TOKENS, +}; + static const std::map PARAKEET_TENSOR_NAMES = { // Encoder pre_encode {PARAKEET_TENSOR_ENC_PRE_OUT_WEIGHT, "encoder.pre_encode.out.weight"}, @@ -186,3 +203,37 @@ static const std::map PARAKEET_TENSOR_INFO = { {PARAKEET_TENSOR_JOINT_NET_WEIGHT, GGML_OP_MUL_MAT}, {PARAKEET_TENSOR_JOINT_NET_BIAS, GGML_OP_ADD}, }; + +static const std::map PARAKEET_HPARAM_NAMES = { + {PARAKEET_HPARAM_N_VOCAB, "n_vocab"}, + {PARAKEET_HPARAM_N_AUDIO_CTX, "n_audio_ctx"}, + {PARAKEET_HPARAM_N_AUDIO_STATE, "n_audio_state"}, + {PARAKEET_HPARAM_N_AUDIO_HEAD, "n_audio_head"}, + {PARAKEET_HPARAM_N_AUDIO_LAYER, "n_audio_layer"}, + {PARAKEET_HPARAM_N_MELS, "n_mels"}, + {PARAKEET_HPARAM_N_FFT, "n_fft"}, + {PARAKEET_HPARAM_SUBSAMPLING_FACTOR, "subsampling_factor"}, + {PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS, "n_subsampling_channels"}, + {PARAKEET_HPARAM_N_CONV_KERNEL, "n_conv_kernel"}, + {PARAKEET_HPARAM_N_PRED_DIM, "n_pred_dim"}, + {PARAKEET_HPARAM_N_PRED_LAYERS, "n_pred_layers"}, + {PARAKEET_HPARAM_N_TDT_DURATIONS, "n_tdt_durations"}, + {PARAKEET_HPARAM_N_MAX_TOKENS, "n_max_tokens"}, +}; + +static const std::map PARAKEET_HPARAM_MODEL_VALUES = { + {PARAKEET_HPARAM_N_VOCAB, 8192}, + {PARAKEET_HPARAM_N_AUDIO_CTX, 5000}, + {PARAKEET_HPARAM_N_AUDIO_STATE, 1024}, + {PARAKEET_HPARAM_N_AUDIO_HEAD, 8}, + {PARAKEET_HPARAM_N_AUDIO_LAYER, 24}, + {PARAKEET_HPARAM_N_MELS, 128}, + {PARAKEET_HPARAM_N_FFT, 512}, + {PARAKEET_HPARAM_SUBSAMPLING_FACTOR, 8}, + {PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS, 256}, + {PARAKEET_HPARAM_N_CONV_KERNEL, 9}, + {PARAKEET_HPARAM_N_PRED_DIM, 640}, + {PARAKEET_HPARAM_N_PRED_LAYERS, 2}, + {PARAKEET_HPARAM_N_TDT_DURATIONS, 5}, + {PARAKEET_HPARAM_N_MAX_TOKENS, 10}, +}; diff --git a/src/parakeet.cpp b/src/parakeet.cpp index b5da73e98..178f049b1 100644 --- a/src/parakeet.cpp +++ b/src/parakeet.cpp @@ -685,6 +685,34 @@ static void read_safe(parakeet_model_loader * loader, T & dest) { BYTESWAP_VALUE(dest); } + +static bool parakeet_validate_hparams(const std::map & hparam_values) { + for (const auto & hparam_expected : PARAKEET_HPARAM_MODEL_VALUES) { + const parakeet_hparam hparam = hparam_expected.first; + const auto hparam_value = hparam_values.find(hparam); + if (hparam_value == hparam_values.end()) { + PARAKEET_LOG_ERROR("%s: missing Parakeet metadata: %s\n", + __func__, PARAKEET_HPARAM_NAMES.at(hparam)); + return false; + } + + const int32_t actual = hparam_value->second; + const int32_t expected = hparam_expected.second; + if(actual <=0 || actual > expected){ + PARAKEET_LOG_ERROR("%s: invalid Parakeet metadata: %s = %d, expected > 0 and <= %d\n. Unsafe parameter loaded. ", + __func__, PARAKEET_HPARAM_NAMES.at(hparam), actual, expected); + return false; + } + if(actual != expected){ + PARAKEET_LOG_WARN("%s: non-standard Parakeet metadata: %s = %d, expected %d\n. Transcription will be affected. ", + __func__, PARAKEET_HPARAM_NAMES.at(hparam), actual, expected); + } + + } + + return true; +} + static bool parakeet_lstm_state_init( struct parakeet_state & pstate, ggml_backend_t backend, @@ -1003,21 +1031,33 @@ static bool parakeet_model_load(struct parakeet_model_loader * loader, parakeet_ //load hparams parakeet_hparams hparams; { - read_safe(loader, hparams.n_vocab); - read_safe(loader, hparams.n_audio_ctx); - read_safe(loader, hparams.n_audio_state); - read_safe(loader, hparams.n_audio_head); - read_safe(loader, hparams.n_audio_layer); - read_safe(loader, hparams.n_mels); + std::maphparam_values; + auto read_hparam = [&] (parakeet_hparam hparam, int32_t &value){ + read_safe(loader, value); + hparam_values[hparam] = value; + }; + read_hparam(PARAKEET_HPARAM_N_VOCAB, hparams.n_vocab); + read_hparam(PARAKEET_HPARAM_N_AUDIO_CTX, hparams.n_audio_ctx); + read_hparam(PARAKEET_HPARAM_N_AUDIO_STATE, hparams.n_audio_state); + read_hparam(PARAKEET_HPARAM_N_AUDIO_HEAD, hparams.n_audio_head); + read_hparam(PARAKEET_HPARAM_N_AUDIO_LAYER, hparams.n_audio_layer); + read_hparam(PARAKEET_HPARAM_N_MELS, hparams.n_mels); + /* + ftype just requires the type check already being done in the loading process. + */ read_safe(loader, hparams.ftype); - read_safe(loader, hparams.n_fft); - read_safe(loader, hparams.subsampling_factor); - read_safe(loader, hparams.n_subsampling_channels); - read_safe(loader, hparams.n_conv_kernel); - read_safe(loader, hparams.n_pred_dim); - read_safe(loader, hparams.n_pred_layers); - read_safe(loader, hparams.n_tdt_durations); - read_safe(loader, hparams.n_max_tokens); + read_hparam(PARAKEET_HPARAM_N_FFT, hparams.n_fft); + read_hparam(PARAKEET_HPARAM_SUBSAMPLING_FACTOR, hparams.subsampling_factor); + read_hparam(PARAKEET_HPARAM_N_SUBSAMPLING_CHANNELS, hparams.n_subsampling_channels); + read_hparam(PARAKEET_HPARAM_N_CONV_KERNEL, hparams.n_conv_kernel); + read_hparam(PARAKEET_HPARAM_N_PRED_DIM, hparams.n_pred_dim); + read_hparam(PARAKEET_HPARAM_N_PRED_LAYERS, hparams.n_pred_layers); + read_hparam(PARAKEET_HPARAM_N_TDT_DURATIONS, hparams.n_tdt_durations); + read_hparam(PARAKEET_HPARAM_N_MAX_TOKENS, hparams.n_max_tokens); + + if(!parakeet_validate_hparams(hparam_values)) { + return false; + } hparams.arch = PARAKEET_ARCH_TDT; wctx.model.hparams = hparams; diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index 74a5b1429..aecc6f3b2 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -126,6 +126,7 @@ target_include_directories(${PARAKEET_TEST} PRIVATE ../include ../ggml/include . target_link_libraries(${PARAKEET_TEST} PRIVATE parakeet common) target_compile_definitions(${PARAKEET_TEST} PRIVATE PARAKEET_MODEL_PATH="${PROJECT_SOURCE_DIR}/models/for-tests-ggml-parakeet-tdt.bin" + PARAKEET_BAD_MODEL_PATH="${PROJECT_SOURCE_DIR}/models/for-tests-ggml-parakeet-tdt-bad-nfft0.bin" SAMPLE_PATH="${PROJECT_SOURCE_DIR}/samples/jfk.wav") add_test(NAME ${PARAKEET_TEST} COMMAND ${PARAKEET_TEST}) set_tests_properties(${PARAKEET_TEST} PROPERTIES LABELS "parakeet;gh") diff --git a/tests/test-parakeet.cpp b/tests/test-parakeet.cpp index 83237c600..58b64835d 100644 --- a/tests/test-parakeet.cpp +++ b/tests/test-parakeet.cpp @@ -59,7 +59,19 @@ void segment_callback(parakeet_context * ctx, parakeet_state * state, int n_new, printf("\n"); } -int main() { +static int test_invalid_model_load(){ + struct parakeet_context_params ctx_params = parakeet_context_default_params(); + struct parakeet_context * pctx = + parakeet_init_from_file_with_params_no_state(PARAKEET_BAD_MODEL_PATH, ctx_params); + if(pctx != nullptr){ + fprintf(stderr, "Expected invalid Parakeet model to fail loading \n"); + parakeet_free(pctx); + return 1; + } + return 0; +} + +static int test_valid_model() { std::string model_path = PARAKEET_MODEL_PATH; std::string sample_path = SAMPLE_PATH; @@ -97,3 +109,16 @@ int main() { printf("\nTest passed: Parakeet model loaded and freed successfully\n"); return 0; } + +int main(){ + if(test_valid_model() != 0){ + return 1; + } + + if(test_invalid_model_load() != 0){ + return 1; + } + + printf("\nTest passed: Parakeet model load tests completed successfully\n"); + return 0; +}