From 69cf5e405f4386a3fbcd3d14f9dd943425edbce7 Mon Sep 17 00:00:00 2001 From: Lixin Gong Date: Tue, 4 Jun 2019 14:48:33 -0700 Subject: [PATCH] initial check-in of the mleap_sql sample code. --- .../spark/mleap_sql/README.md | 10 +- .../spark/mleap_sql/jars/JavaTestPackage.jar | Bin 0 -> 12709 bytes .../jars/mssql_java_lang_extension.jar | Bin 0 -> 4748 bytes .../spark/mleap_sql/mleap_sql_test/cleanup.sh | 6 + .../mleap_sql/mleap_sql_test/mleap_pyspark.py | 237 +++++++++++++ .../mleap_sql_test/mleap_sql_tests.py | 164 +++++++++ .../spark/mleap_sql/mleap_sql_test/setup.sh | 17 + .../spark/mleap_sql/mleap_sql_test/test.sh | 6 + .../spark/mleap_sql/mssql-mleap-app/Makefile | 16 + .../spark/mleap_sql/mssql-mleap-app/build.sbt | 23 ++ .../lib/mssql_java_lang_extension.jar | Bin 0 -> 4748 bytes .../mssql-mleap-app/project/build.properties | 1 + .../mssql-mleap-app/project/plugins.sbt | 1 + .../sqlserver/mleap/PrimitiveDataset.java | 85 +++++ .../com/microsoft/sqlserver/mleap/Scorer.java | 314 ++++++++++++++++++ .../main/resources/adult_census_income.csv | 4 + .../microsoft/sqlserver/mleap/Predictor.scala | 66 ++++ .../com/microsoft/sqlserver/mleap/Score.scala | 235 +++++++++++++ .../microsoft/sqlserver/mleap/ScorerTest.java | 98 ++++++ .../sqlserver/mleap/PredictorTest.scala | 112 +++++++ 20 files changed, 1394 insertions(+), 1 deletion(-) create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/jars/JavaTestPackage.jar create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/jars/mssql_java_lang_extension.jar create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/cleanup.sh create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_pyspark.py create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_sql_tests.py create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/setup.sh create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/test.sh create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/Makefile create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/build.sbt create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/lib/mssql_java_lang_extension.jar create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/build.properties create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/plugins.sbt create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/PrimitiveDataset.java create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/Scorer.java create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/resources/adult_census_income.csv create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Predictor.scala create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Score.scala create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/java/com/microsoft/sqlserver/mleap/ScorerTest.java create mode 100644 samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/scala/com/microsoft/sqlserver/mleap/PredictorTest.scala diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/README.md b/samples/features/sql-big-data-cluster/spark/mleap_sql/README.md index 5eb8f803..80497f51 100644 --- a/samples/features/sql-big-data-cluster/spark/mleap_sql/README.md +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/README.md @@ -1,6 +1,14 @@ # MLeap on SQL Server Big Data cluster -This folder shows how we can build a model with Spark ML and then score the model in SQL Server with its [Java Language Extension](https://docs.microsoft.com/en-us/sql/language-extensions/language-extensions-overview?view=sqlallproducts-allversions) +This folder shows how we can build a model with [Spark ML](https://spark.apache.org/docs/latest/ml-guide.html), export the model to [MLeap](https://github.com/combust/mleap), and score the model in SQL Server with its [Java Language Extension](https://docs.microsoft.com/en-us/sql/language-extensions/language-extensions-overview?view=sqlallproducts-allversions) ## Model training with Spark ML +In this sample code, AdultCensusIncome.csv is used to build a Spark ML pipeline model. We can [download the dataset from internet](mleap_sql_test/setup.sh#L11) and [put it on HDFS on a SQL BDC cluster](mleap_sql_test/setup.sh#L12) so that it can be accessed by Spark. + +The data is first [read into Spark](mleap_sql_test/mleap_pyspark.py#L25) and [split into training and testing datasets](mleap_sql_test/mleap_pyspark.py#L64). We then [train a pipeline mode with the training data](mleap_sql_test/mleap_pyspark.py#L87) and [export the model to a mleap bundle](mleap_sql_test/mleap_pyspark.py#L204). ## Model scoring with SQL Server +Now that we have the Spark ML pipeline model in a common serialization [MLeap bundle](http://mleap-docs.combust.ml/core-concepts/mleap-bundles.html) format, we can score the model in Java without the presence of Spark. + +In order to score the model in SQL Server with its [Java Language Extension](https://docs.microsoft.com/en-us/sql/language-extensions/language-extensions-overview?view=sqlallproducts-allversions), we need first build a Java application that can load the model into Java and score it. The [mssql-mleap-app folder](mssql-mleap-app/build.sbt) shows how that can be done. + +Then in T-SQL we can [call the Java application and score the model with some database table](mleap_sql_test/mleap_sql_tests.py#L101). diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/jars/JavaTestPackage.jar b/samples/features/sql-big-data-cluster/spark/mleap_sql/jars/JavaTestPackage.jar new file mode 100644 index 0000000000000000000000000000000000000000..296a16de73455f8c967c663e84315701cb0fdc1b GIT binary patch literal 12709 zcma)iWmucr)-6)pJ;B}GwJj1T?ykYzT?>Wa?(Xg`#e-Aar8q@Pp-`Zu&`Zxb-`)G$ z=bU}-%%8k@R>rgP=AC2CHRe!}hkt<&gN%#}Q)Usb1oNlBhj{^`2+|N|lU0)DP!v~^ zl?JJ6uq#SGO~AlhD}I<%kY{6`!;oiVp8hb?tj4*>vvc_M1KXJFl;RYNEGsPPDa1VO zphSFsV`5OKODhv!V%)jjy`+0>N*b}T_;pVOa zHnRg;S#bPb{hxmw%YRx8c5~CPc6IWwvQ`1Rf*ss6HKYaE&Fr5)Zq?G)#?!`s!bHiW zjUfuDP%TlfHyAQ$O$#x>QIUpi6z`thNFardnbUvB5pSt=sBOKwEmVGjqtJh!23FN} z;s`jF4|*LWw^H@dD#cNt#a(-N15bn%c0%mw z2@I;B{0dCopxgS9@e4auchHqDzS9ucgg43t0cb@W zFJ~=20@q|S6e7T5i|?zKKf(-!v~Y4an=`(Na4nBw>?kwQ3LeIm&lcLi?v+PjedItP zHICn89`8J{6Q%L}k;kt{XXXOEMMIxLDBsbC>>T;4UV?lkLFyl3Na_-;#$IY>{dgrS zgG-a==`R3P&PfCH_70Czo7IXfkrnjUvmWA!O=ni?z*#Z1xY6 z_A*mgfd~s?VD3^g-GCxQcty+TGS;|ykmqmen~4HWQt@iF*wVFVIIo_!SHKcE2K-M$gd;vRo0v6 zFeTrZuxneLWgW^b9ZU(Bd7HT>yemIP?a}e%)YVV=+S39-f#sprfHvn~Dy1MMCtcdX zu%D}vIY6sY)`Y^$%=Nk+8o_LBZ@fXT*suFAV z+Cas;;>Gw^5*Z5gBNC&7NYU2Eb&AXEV6=~eiRR`gvFLZRR5JR&z`-`h(%NFK9F4RF z!|7z#mUmzrkuRZZkE4Aa{*wH5a5D}EB!EL9|oZ>y4A=jBeB@W-@65R~TT7P}bN z&*+}ZqGcFSAyCvjb>~pMvYm>MqhK0ALO@_}(y<<*pE2+pAf9YpmE^_X^GboHi)WI# zk4ECvjwkm|H129jy30jO6dN8s*{Vez*CJK*5MJODaNKWoduy>#7b7fcmF$UJ7Dby8 zNKZTeS#f-w`gT+1jOM&2XSF@z@s-xv2WAoqR+^>Q*Z~1+Hj8yxCeOJqw6fuI@2-!B z>2AQN4E?gF*iio;gguc9H?zmEuBZ}0Cd_TgAHKbij^CGqfxRLkauLFPK>J+>)G`va zgArh0j-Cs^e`j{K{~u*)jORb5H>!_F0pPxFAD3&PV}#VzFoCVtE3x=YI%#y zvZJP%wL7OT$@#zTZ$)5`oLS)i;bvT7tCT-u!SOWFzln^@M7LHDfSmF95}M@mK%Romekq- zemGE7pJV!wMs*J~#t%^KbX~^ZmxS1Ke*LhqoTXSx{00k79k>!OT|wxv2oRV+y;=_Ul)~d%3*Ko zO;NQq6!xP+fCTls9?an~X0!X*T5-vPtxZNP=yxRwJ@tm_SBP7*IX1NJYBQ-QPF#yN zZN>olECyL8sP<)Cl$T3tf90N-0ZnL?Hw6W@f~J4O8MEU~Sx7u!nT;~`vbA`d@B&@Z z!9vSpD!nzYULD^apHOrG3xkD~ENt6a$OZfTU0iH)V8HT@s(XCp2@-^g%2ry;1PIXq$A$8A2O(uHIpYx_C|Hi)#b)q!aWr^{L6~P_Uh5 zks}AbWxs)?e9h3JhJ~KmbaV&dk>?0r4zs9~kEu0dG!GeygMNaPBmNuHAx6T^Tr+F^ z!b%Q-Za4>1dxeuAFH!=fObyVn+7y04(WWkbKO~h`eI~*+a+i|f+Y&%)ftD}hdTp)Y zpd*WK4S5v9DWOW(du@6F-V!Ohl&P|%lI{(10!|@m-rkEmMmA~MAL%bTKFEv!I}Pg1 zz9T&F5_d4k3Job#fm5+; z%_eNSe#ot~gLI2qk|kX#7xO};i6a<87_`Tw#xk#{mlvvF^&=)4h)$x0@Q~+(8jeV$ zzR-9&5?)CeqS&%LGHA`s>I|L^5$oBPf>7+_D~xBcv1$~O0@&%cPR~ZJ^Z>N;olZRXvW7#i2;tqigo;)qWvSDrKPo0IU1iDWG|9kTQk9g z`gLJzx#4{+0gBvqa8qn0B&%|ISt`GYUWxyF5c8Z~MD0MEM`^BByy@s<$8t)ttAKH} ze&=n-J6>*taC}n4+x<6`_dwU8fS0QxECX?>*0`}O_dW3~^6`~mEuBD6=+j%0@k?cZ zF}@mV-a5VHx|C!R_qWhZYu*&eEc9P*9opyC*H_za?^G&6Kkr_~ud33(6Usp-Oh^KF zYGcROaXq)h&jQYc(^ti!7w=eH8%~$rU&Gpsp6=XPvtcIfhL7*l98;ajatFPb*@F{y z!7W() z{q^6{J@tQ}dnG3|3wIA!M=7v7_@5l#qj9E$CyxIFA#7uttThM=0fF=#&k#6Nq&np= z;2J}R@)thq5VgIx@Jid!uSB7n8$!Ay@F&{OGyb#Y(&KHq%q2^bC0gsp_ znDwUC104{3)tkLGhzOU$PMz?kx;%3fFFl{aF!r$biv_u9DbDH=oPhv3ilIyFcxu1` z=jd#kwNRRYDRSxj1NyFJJ0+m2ja9}TskHm@)_aRyQKyp&y{rxMo=d>~0zWqAhP6Cn z+>$^Ik6{p}I#;I4=FC(Q=YR^ugln5K#gb=basof+#@mfSnqkAWXTXj zsgSALbt}iWK)8uAf!>%y;$61dut;`i|K15<$@p82=A!VA*%@cH-IeO=U0-4HeH7?m zpWB76)_Br#V}*9XZe^TR#!BPcm^hF*ujoo+C=qW#5WhQJ!+FRhplj&`gWI5X&RT0& zNnE;4nVmB33Ajx~y@s=eA-AAKM%cJgi>n^2I4aB!z^#uAP`-Cf&kpZqjF}`e1K@rm zo{8Qf%ELuxm9IagZr99ycxfo5I_6+hv~^cHuWc0!ZW+k%m2KbEl{@Vy2h-b_zWmihi}qIK{XB?I<%!`K~|<7VUG6TezBo7<&TY0I?9+kGJTd+ zu|iPFx_3oSC5=r*37ujV{7&*aLufG*yyxLyU_=oAt|4Oh4;iBJ*Q4&gS|M#ecYG^? zUjb{Q_Vkh{Xejn5?y9KSp$%-MiS!YO66TU`)c7%Q<<89S%y`VbEHg6?!46@1`nD7= zF9}$d2|WT#^ch*z^iXZ9+ia`9ExCz~28res|MC`n+O@HZzRS-1X#XZo<}f42e`dR9 zyXWrV>1^C*%k7a~a`K zc-)t)yO@OsG~8)BN?cv?gEAJZP;o9H+FeKB)O|+ACe7|Xx7LmkSEuqI80fOeL7$&+ z`8Gz=TSb8!cM6L8lIe%C?0riN-+fF5Y2xJ?_DhChy!`Y_1jvq_;>x)Pdauqk=ewYH ziJ#D$A9!=`h+K;|N!4VUTlR=JA>SK&Kb{)#KTgy! zo){Z;)0^OBHA^RJzcpubyg0Np7(-R<-Tu^73iS1#U3AT={e1Ybs#b6F!u=bOJsSxF z&4`b0vTdw`>UHjvc2O!=YhFuzzhulM!oOTT{St(Nhub@(lD){?WlN;sx^w<(IJk9z zM!S`pvn#v&eA+|%v@B)9Rg>cYNO%prWvW37l?vet^a1#U=GP_7wxOZ!>8i|sV#1%+&Mz#= ztAG%{Gl@?2>wP?K9P>RYx#p#&XxWa7DE)3mNd4|v3Uaz%UvGiM!lrozQ6{dUtNgKo zyO$J!KRn1(CnY70Vg6!J0C-yDOC=y>m!|WBkUaI$0bzTiN}U$KlY_BUE$g0WHDXH93sa(whXONfx+ z2`?MydQZAzYKoC5@`9#Hib)#b8D0$HyvSb!ku0-bkHwBuI#oR8pNT*=IIK}E*_TB6 zIvCRIEj=J4TD>EblW!{Du3BKPaVg5q*b(t1JuG+XrDVo;_@62RXx=Upx@Q=Ymj(mz zQaXV#`6^Rco1e#tt7>=Xx@&gI@7Gd#4l=;=TQgTtA=+1H2z^6#8f47$3m=ok`UdNP zpM;`?mdK;)SGYflmL4wFglr?r1V(RjbyrV(w@&EZE7f!CEmeUC@aNDbE6# z>0UtLE)J#XbRnHKLRw8jDU3n~g|_;osKts;jl#^ zQ);DLRaOlWYqomhSDY&Vv4KYxm8)jEERD8X30K9ZYXfDTmZ{5^iUOOHnz;5+bWa~W z$M`o(T_+U9GT~B)npqJVz-7smKf^lS$Q>ztN-@*vrFRL@>dYNs+Bnq5;|ow|{ZXkP z6n}%*&Wi>0^h3V!B!j&WWjdPE|0-T?W3gs}j>|-l+^w!Iqn7K>b8J=&FWkMNbBE ziPbQayPS7ZNJ_tc@?X!%BpReqwVN38Au#Vm6>4~!s@9yReZ$dsrIs<|{#djcH#=_oE zp=5@!xK!=p)EjPlDXgimWrmXj>~HUmb{JB31C8GYfl%B;;1c_FKNNe30Z{Tl@X#0Q zLlud;aG$*1usL8nMYt294T(}qYxDS?5^n8V4a?WJh=GLda+39r?M=9HT`^7>jocbi z&UKU=&{yV=?&u}#>EHjDgJ5C1;@Gm|yiRJ6m5`9&XjK@rL<#7}Slwg1U|>LG9_U?0 z^JEL7-{gI_y|QppZ7o&bz_J%jg&fJdB=eJ_sZTR3C8WG1GJ(B}UP8blHbbh?MUwc+ zcqj>FNH2SxHord9>mffC)ZmZv3xSmheDc~(owi@DqiswH?Xn4l7q8>xL18JXuP(}D zyYI11;zx$rDe@m|Sccl-0_n5d$A$XQ_U~WUU(PytTL&R1dad#w6lxRdR=*wP(a0rlpa@T#OSl2sz&?7*P3crO<;b- z_W~lTnV?b2zE`3`klD{ob&W-IM@pTVOe5=Jk0L<0d%t}l?v+ybnZWuzQh+i}JL+*3 zN)7&*oiH3NPK7Z&zmsd}vS6ANoCK|ogbdEq(QKz!_bvL?!%y_Gh9)2bbN3^fa79nz z*Wz67E=Zsg^(S|=Pc_3@S6HkTA`Ga+p>0gcSZZ2iWglsV_KJs2M-%(;hR$)XZmH*P zEvV0{hx)gZZzvLnjJZMuIAgpTV#imf2La&dhIeeE4)%PGR)#9N9P{=tIUhsNHzK+C z5o0p}Ru735;RVf5&Q+X)HPwuB0IMe{!Y%J@)Aj18Q9s!^-{!c{sNp%5;Mh~B<-x>6Dj}>`aT575c-^{I^4$(%6DZNx} z42$z`xHbOWvFN}Zj~hOV_lH=2*Re4Eea9kYVP@k1w*OblqN%S<@~qxq(kkF0A_u34 z%-h;YP@44=z=7lol8D)}B9?tc6qKMToh3Y0I-lSF>Lz$9CXAfbT1UZ@7=$5>TEA1J*IhMaCmZD> z2~|UtWu?6=n|GqSY@44%sbTEN8sw`f#fFC9eqp?-gPh6Lv#6?>N=7QuInpMqT#`#} zR~_oP#uf^sZ(BvAc6XTos&wjBQF7sTSMGxq(=t%a0F#;*U(Fm>Hyr?)C;EbqaOd3p{dgcI?L4R~}JYF0lWL$X3O#E{B3dQl{9>3kg zA2I_(EDor!Exs*af+I8DktZ`|8gMCqN*RXtlR@MKn40L}I%UG61VUZ%>>w_4;G>kf zfL)g1{)n2wa>G=6xpHy-(9K=UDM#M)S3&LoOC{_{N1hb2PeliwdS&|a=d$yTjPET= zXK2NcOPr7SFC~08njHiu)fZgps8BDkO(~rC!}P6m&RxlJIJ$Y*q$ULhN*myv@B~gP!~%;C=3nvyXL+e(dyB zq3L|+FNL2Pbw34mzXuqU1VxiqD5ER%9MBpUEFtOwPfotem|Qo8ALq^t=|#Mf+8%#* zOY5b9roDwAf{LLtL1UiH>jMs-_PYuZht>*xVFi>KTE2*8FD)q!koMV=Q=e229 z*uV#Du@AG7XeWe_{ELLBDpGOSVsjrHjyb);Ks1bi5&;St>4Z$O5E9*}Z|uX(SN7*3 z_(aWD68u|(*2X35Qm!>&S9qgRXvxN-Yj;6Uz4!32-=sDO=9FF0Q4~}}k^u}^mM4Ch z-3Yc-Z_cVUM)h=~E&i|qj4F~lNY3#shfH@jY1ssaX zn_Nf>Kc=d2Qp<1(FD$MU%It(oQhc~S7pkvI37%w-OpsJxTW;h+dgjyHr=+K-7QZI4 zq`eb`>`E*D>W!X&+qIAHA0LM~`d~XC;li#JwjF9iylK1!OWmvmP8q0Q;x?B0r&bH5Q7d#aw%nR>T$XszE_WV8CbY5#GK+0fo%J~5nZJ8<%Q6+xA7gz2p&^(%45{NBhbG8%AfJJH?ajO|-ZGza0{)_sw_EBnN zGj&>hdwkmdMv+AKSMh2;@$zpDp(o}?-yAT6&u+l5A6ULD4R+WpM!UZg-cl0lV={%H zcu2~}y}>W*={4YPcxPzq82K9Ov!!<9D5oCkLV;43RAuj=A%-!Dv&xF|_G@@7z;vF} znxH30AD0$7NN96Lo#^MD>@BbJ?sUD}`v)ecLbOq=5n(OW9xc?8zo#p?!tSNGx}?6u z#pPRHLEXBXUJ3N;04Wjs@+R%R51nV0n66#3Uh9g7uj?bVb9!Ywkz z&dX0H#qT09w}c}LtXLH_T+(>I6iE4n5hrLgT2+fAa|R6_cxC8Jt49E-ec~w7ddshc z@FPCr%ZBkqC&swNePzRjs!9;kO!Dp;(N3bf2348!c31BvK@ju?cn$b3$ur%&^C>4e z2FmFs=rU)|H;~a%6f>X8)=7o zt~jM8s&n&j|Iqbrr#T)b!|*BwK*d@+>!qkeTxj`9iXu#SV^B&Npc!iG z==;I*SjFCWcSo)-F%cA$j3qd2E~y>6TimYm4o#-0ejuEZ$-`fX^P z{MA9(RBlmba^4~S!di`l?cSRMcv+cO^vhoz$60lKX1;yAo4>)bH}(!63R69U`R>?l*<`yBr~G=lApb zKA9p7Zc99q7%Q$GEp_GlD0$o;p8Q~m^G>|ufMm^Y&N3wRplD>*-m+e2rVoVhlkCXx zQR?X&d3e2hA1p2Rg^j?$mc$J6kM@O@I_Ad_BH}jD_PyL}U};I|QEn{&WDENt?edU* z=F4>v=L}yASY=UOwAM%8TaK%NwI{rJG0lE(4iF_#ca7`O%fS`sbBPYc7^zXmH1GMz z+=0lG6)a6$5c`@R^eqQI;Ifo)4(b8#oAJmeJgLd%MazY1L_OkZ#Ew?p&$)JRXr$D= zmxfMO z%x4qIY%L<~XI2YJ0L))G&jQ;Y`dY0Rn*v^FEp*`^_?rEowMO`R&6Mr_J;BrCs-D-$VUC6G$R zXt4q4XU1$i%)Xscr$mE!QSc=@P&+jdMuXsg*S7}mp32{{-ovx_{om)D|4G(U@^G+l zwJ}rna96fe13Ow-{3}hgPJZ{uzr>mp0EodUv(g0^4uxR5Du+fm(4-p4gc^w)$UvAj zo|7)KgahB86=j4n4Jpg247R24M_2Cyi*1~)}CZ7L3fb(<9P%h z)VQWqxJpj^X_IsIMlys`DWlXQxa=`8bG42ZRV0@73nMJ7eY;sm3A1qXkLF@PMMUo; z=QOxe7K@}O6cqKTt`!CkMo0xDB(y8O`b3|VrOiu@qO(=%u9JSzW?MSY(%uRjna8C< zHBB3htN8SxA?w1n7pEG%6o;*@T5F^Lx$n`K`rHU>ig@C9(T2RW%h!*{*$|q0;UPF1ESXiB*Asp@Qe~^82mBmx zb~B=>)Z;1C)FR>|^{5RU1suF#Hk0}eR=m-|@E{gwIkp}l&V9)av&wcua6bGJ5(K+a z$h5${SV16^9I?s`KUo9&UPv7hGJ5qw@+gQSIFx}vG|UZk+Dvs|K_}%hCdBHBG&k<7 z@NCLD62Q>SGJJ(1)xl&Bc@Tj|%^O8qOoinZCW$*l;owd|xKVJ8u`Z+R33bCY(z--t zXFMJtLwz0*s2HLf2@jS`I44h)q)3>d**G$L6Vx+xM63)03p|0^ctG0_0sQXBc0gPG z2&odSw}B!l_nrpUN8*h5uPvY}?v~JR?jxurnv#mH(ImkHlELO9EpY0h3H)kqdNHF? zOYEE)wOrNuqSGip+OIIF8rQwPZ{eVRVfM!QqJHk9W@0ian*g% zmO2nnj@}R0b}3<1+p;bwpg_VaEmLC3chDPZoAq5?M1oPRWZ=^`tJuQ8wW3&`!ZZ8B z?yUBV`DF__)P{;>;&W=Wvf)Zq(n|D?$GzP`Y;?kg?PJv1$>wvlWeW!7lchalseLam zCCa@bCmDZg+sCI)Bp*L9?q$NrMH_~|+z?zVC@ZECM}I5D%`AEOfsIePbG==uF?~ER zi&?f;rLUV<(t%S{_h%9=YPLWJ2?|dYxum-DW9n*mp47M+{ViAHw)QEKm9vT<{SnJei4a_Rd+dh9+Tw^b~9Ap%v(85UT^c!VZU- zFp2GUiW)q z)tW8JBh=NfRBb&KZ!4}5XW_>?^mLy$(#b9Cpv%H)X*t~aY*RS{Rxe-p4qwGU8faCr z(E5U*^Q%YKqhomci00bCm?>KAV;$?yz`H{NEPV$N5A%BR!JJcnHr%H($Hv3KGW$;Guj){-xs-cq887 z6;%RAQu#5xtzFaJOWmC+1&rmdQR>U>!JsN;(g$VORiWxK`cdx1KtaY*Y#=Jt*An^} zuQgie98g#DV~b~gId51A$Hk0L@pGH#+Ia3BFSY_FToOK+rECi#I-`cWr1{$7fPJ;&)~y$d4v0O) zlQ{cnNz3U`D>-!kp(g`5IJ^6(S-5%F zyQ?{Qx&2RHw&qF#!WVOpCzmDn0m{VTnHRl|u@H*W34xsi;I0jrb%RkPS2v2K6-1-C z7tHpMw(jAd-VzxFS22RrQn7Z>?g`1^+GW0~HAXSNDP+zS33&gdQ~&PqaU41dQ@3A@ z{zGM+GkXMNu*8`@z?ZVLXMs>(a>Z=4;4JozSB|N>?!X__iSfH7suNRZf;d7nkXA+Q zvMq!T;R~a;_5cQgfO^KViUmNYMNFTa>xRh66NszXM+JbaI4IjoU*sX5`p-})N|2ce z8XNRMvAoxIjhYRTgFo}JlvT3wxhU%-mK49jUt!y3d*4QIU-+w)HQXy#cKD6f(g4b|&V#~)RLme+hjIn+^XL&~tiH9K+KCU%4cz8-N%{{`_zY4@ z7H?M0%SJ8mRc8AhIm9{gYL}U&mbb{j;}i4EWo5c?J4TD9V~al-8;02QThBM6Hj=_c zsAEVw_k&1`)XjzSB=dHS1-l!GCf@2dd9G9rTWQ6$Y3;>51jkv_aI*fqqvhS4zarQ2 z$k181db!M7bGF(~Z`fHb@6FSc3qb}*FXpl?N>vC>9fpO(vDvE|qe*o+@e4;`0T9uQ z4whw!n3piM_!iw4uL|P~{Y&=EHlMYbmrGZfAu|O2!5l#L+Qcti$cOwN$+qwh%jG6G zsk%UdbrW^SDJCZLVkOBt)tFf0ukBDf#|@YA_dcwh4RHDlS^L($@kGNX@%j*BBt1g` zbO1G#FO4x3Lsd}Ijk+o3qNhEM?Tz_{lZO~P_Wc_N% zO$wA5=3IvI5LA?EVG?kzO!CBLQ89y;yr@^Iep$*V<}j(ic+gaYaUSJjGR zrxQC}YPK*#Exdfa_o^&izEw#36wj&JBaf44SV7Q0%_|xujS-W)9S;S8D4gGJkQr=R z6`p6m#=bCGa+@+i<8N9LyW@k4u|5;6mrvuy(Z~wEAr+pwXGKTHS>gO_BvClo5f!+O zRQD9MNK@kU>h%b~UNXP5FOhDnoy#13VU#wYg=>D~I5J#;v@}>ZCh-%uQ9-NeBPKE8 z8+I{?5XM-=B9d(7n;8bTaf3Y>g_*;o=1y~xBhwGVg5>+(`RRL=oARH2m|a+z9!EbK z#H@`Z4!$@{V8|H~9S2S9gVx83zMx3+8Gz4~Q#c7odZnloDbgLQ^t*WJYz>YRhY3vc zsTH~S^Zq$Q{`~0wRQ`M#Gx>*@0t<%^^M_&j&&KQDcJTil|IHk(A`c6P3iIDh+JB%l z{0Bby{p~;R;m@YO-=^(9V(aWa+@I^+|GuLIH2)6w uUyfCO1^sg>{VtV%1m5$l`0L;F9}B07Ji_zC2Lprl{0e!_N(}bjSN{iBqjJaq literal 0 HcmV?d00001 diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/jars/mssql_java_lang_extension.jar b/samples/features/sql-big-data-cluster/spark/mleap_sql/jars/mssql_java_lang_extension.jar new file mode 100644 index 0000000000000000000000000000000000000000..85cb713358b76adf4aed4542ceb5329840bd4508 GIT binary patch literal 4748 zcmb_gc{r5q9v*`%*~xHHwn!PIv4w0I+Zg*YCCix5V8|{@_GPS5_Ut4vWLF4TDoa9T z8EdwXWlUrchx46tRpLVGbt;BLf_c$th?7 z0BXRWrc7~|fpW%({W^o;-x(M*%GKG;+1C9xxg3AWb#rlavvKvbas5rz`QKF8BR!Fh zNGCfRZ+9CfH?*_Uzi|F~`IY`2jw>32c1L^Ks3P5wZZ__sC`Y858!nmT%L?V{oP@{I z^IjIZ4;0~0s=%q#>(dHhibKspcyLts8=4@gh80_K94)QT6qf4rR8Z6c@lm|D9#gqG zcK{)NgOEASm4P`w=Zz4onw)_^*1T+M5B7G6&Vb~Z5~ge7XWD3a19LU&SvmWM^Ik{I z=iFq#+RB7p0dyg)2P@ygp}b z9wZWi0Zl?P0-R#q-q6mHi#pxj@XqCn6L$#@V<5_n0y7xSHI38<7d9P7haVQaj|qJy zvYQm5CE?)U<6d%cm=|eq9mW}zSz|;a3~O>2Gc^}C=ubtKoO{2tp8B!2kh3_><5ZZk zXB=IWvWYoJ(3$0dPfAsr0qZq22>pu}W#_W=+*^~!hk7=3Jbe%DNFwDz-p#YUUXmqh zT)Txauq=O7YLU;G-lU@fb-;_#B$?Y9oJhA1eRc+Q;d zo=OF+2rY$0Emcr-C|?gcVDkfDHOS!#N-&e*|Fo(0bE*Klz^KX)0Ni^ZGL#$ z&>>-9F(f!-3^R>)Ixq0{^pL5#WnEe!pM{;5xVVDVH32NaI!{<_VnlTTQMgf&nVyT_ zCn3@*{4+_ny@WlH+h|HHTa9>HUG+NzaI-xc>vzMA_abhXykAG;eXmNviE&M8L+kjs zxW=a&J~pnOGsY5=ZqCa1Tvx{)i{oGY@YXa=HeIX3J-ESKxA@v%SkYJ~OKzdHarE2g zI$j$VxQFbVloexGQbD>g-4_R_mbDU`Yn@x5BIz7T$jKHLI+~BgjMcXd45Z^u8C|-c zsJ{XQa(2>|&fW%Rt;K5D24={{gxzG!hfC=9gV}po`oXNd4Kl(r*E!av{b`8uUxUg@ z&Ot4m3+rjyCN4bCqrxosNs=~~dhC4+YU*F?_3-a%3xBjwNI-bq*k97%;GVkD|DdRz z8Cqu~aEV1fq9!JjW`^Gn%)YM40a{0MWFs*gXbz-boY>@F-SArUeG$Vw&55}39GGa`UjTeH1V$^ z!HN2~kjE!ZP>5&AI+ldj=jQDLwT9&hvzy;I@dhlOej4#|;wLxCfN)zKlkw+$}_u8NhY$e*Cszjb*8%R2DjDRZ)lEQ&B4qc zm&4V3hE5`!l(J%?TEQSc-8542)m8+nE<(*WG2$_wK&Y9X)A(Fva5{30g%2fqa05MO zsqucV4|0sgs{+(A?g;Xk09C~9On`PLQoJWpD$VC_TgHuiCdO2mL%R(hSuuL&__G=8 z8dGe=tjTrmn=O2-xN@KuDlr@K}(oPA%oB2da$o z7d9+P-@s@2_&YvE-~&``h&G{RR;~a$@q_a(e9Q8}$=^x6Ggt8^3TUHmS}m|VjI$nH zBW;i9=%6S)2Mq+kULcE-4hK}Bw2G&;G36js{j@mo(%q4XI3MtoBmS)>>2!ggCi@Ii zbic9~)T{?v-a0z=Wr1M&iGyYEY49S;b5RZb6>rr&*|C&p%dS@7UgpGveAL6em~tyM zK!DY}SJ>EH)*3G3(zMXs3Wbzrz7$1Bw&lv}Y!8}V%U!0gEWFjLT->~tJ+3~#IG z@y=wMfMVHVv;n=3=lb55yho%_^gFs@smqKNcSsmSYk(5!l0CEGGQ}xA+tgI%YhIOK zqyxRvBimHDay>wvlhhxDm|sT9=q&S(U6c-#QEu1d)9D}W9+!HPKHjC!|)(>e_gpWj|2BUVNi7BKBZvQ&ko zEzbd&vj-BVY!u1buPfYXPf_)QMt8YMlC{epj-*xnV6d)Bi9@PA3p46|y6ueBH-{CO z+gA=@^!Y>wT?5r`10&@QTQoBfgniKq>`ajsSvRJy8f+vfO209^qf8LC+>>v}IMBV5 zb?{Mo(|K)S*gg2z0r`;}U6CkvLl;NG zfAKr?KLic+wn2HgJG=hDr%i@%=RqZ=Ot|(~=ojJK>T}f|$0N^BfzLP@^z+TtmZi}Zi@G&=L8VcO zc{5&6qU)lbBJwJmAjOAa8aGzX&JWld*S&eY5T#r}v9F7y;#o*BTIW+8Nk);P^c+D3 z3wKW&4i@>{7`?c){Ycr$O=U^I^}KN1vWxC2^D7ev1~S|U2Jwbrg3PjDYhVYJ`xHve z@B*;p{j%{WrsG|oFz=V@QJ1zBI(caZ=(ckD$5$Y*f^9`{eN>m=U|Mg`U2%^>|5^OG zdFyqT@>pUp_KP2D!%lm+Sn>9Dx?1e0b~n>2{6wT-TDweJ-vOj;brZ9nut?94CVbKt zl5OPWuXdY-S!fu^((w6d(LSNA*$qeFylg-j{e+?$ni?YSDvI~-hO$7B;E^Ywiqx+A z3zZrVzwUmSTG4I~~TEf%i zZ7oTJNj|wKV?NzjyEZdw)8;h0L9@)3DuxzAlP_ycv!n8jPQzR-tjBJ1J~zgb?A@+d zd{S3x6QwZ9LcmyMRzN<57MIr)D|PRMr=ALy?YYX9zouNs$;tAmA*2YVKDd7eYL>(% z3B7ZCCG-BRuhK_bwn|Lt2?PiLkPh|fku?45Evx^3+Oj_>z}Tb}FIFg9XZDI&=de^E z<$0dO+Qx#MSiU@9M#ODiWL=(llp2Qc%6%I=^&_ft+}GwW7%01n+uEmaElP!sDsn=E#ob*U8#X#Ontz z+1e(4vtxp2hq+fR5=lteB7Jf`JCmNS1ch}jC~cX8@SgnxTB$r|EgZJ`G)F>z%Q3JXGLD(7%nyjJppRI{zY0iZUG z#%-ePi(ahSOpPwO+JKsoTf{#%JKwY+$bDqRThh4l3y*mK0<)V{N}6MBC~X|-!u!n# ztLt)4SK{_0n*#p9JslNv@*9y$W<{q2>QS;|0lPg=bSlHy{_z}Dh_A;h`Ku+F9@iBe zB+Y8=&qqitl}Z%DTxKrh1UdFV!LFELC2$>^G;4P$-43>}OcmF$z*)^x2sFmr8!@d?87vO`LAqwTFo~~m2Q^&)B|AuXWF_US@)zTTema!bGPp4P)4H%Y)T*Gb#oMXT>Hj~iXhIcT_Ah~ zWb-yfIWSLrBHW{ER`NL|&9{e$&$ix%4se-VRp?kx(Qp=A51#fgbP`x!6K`DYdAU6s zKzueJIm^GN8m|l4`^RF@cFf$Q-8R0Uj&P~_@QwL^+;;Gs9moFZMRUK`8&!18Iy)zx z+mh(oKpJb)cS5<>2yS$ieExyXh<&;?=Y+Buv*c{MZS#S(^}9`zl;MgzB32=DeHZt5 z(aag(fVk=Ggyi+h-Pv~Dnmbos_=#R@uG0I>m2PN5{yALh4Vi$N_+Sx}#%+49Euf?i_#g$4AEf=lGS4KSH4e z{Mt)@N2&Vv;q<$TBX9jZ-XA_5dGQ~e_OH7gdF}5(cIfhduKm$@{~G1Udw&n<-=O@J z6aRTdf7BUD-$4(b0{$NEUpn(2|KE{z`H>=yw9EHsI=tuj)mr@dEB;p{^WXh{+*?OP h^gV7J|Nr#=iI((Wl!p@T=njV((0s^@b&N-&e*y?;?Gyk2 literal 0 HcmV?d00001 diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/cleanup.sh b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/cleanup.sh new file mode 100644 index 00000000..fb8dcd10 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/cleanup.sh @@ -0,0 +1,6 @@ +#!/bin/bash + +echo "Cleaning up mleap_sql tests" + +hadoop fs -rm /user/root/AdultCensusIncome.csv +rm AdultCensusIncome.csv diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_pyspark.py b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_pyspark.py new file mode 100644 index 00000000..c450d166 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_pyspark.py @@ -0,0 +1,237 @@ +## train a pyspark model and export it as a mleap bundle +import os + +# parse command line arguments +import argparse +parser = argparse.ArgumentParser(description = 'train pyspark model and export mleap bundle') +parser.add_argument('hdfs_path', nargs='?', default = "/spark_ml", type = str) +parser.add_argument('model_name_export', nargs='?', default = "adult_census_pipeline.zip", type = str) +args = parser.parse_args() + +hdfs_path = args.hdfs_path +model_name_export = args.model_name_export + +# create spark session (needed only if this file is submitted as a spark jobs) +from pyspark.sql import SparkSession + +spark = SparkSession\ + .builder\ + .appName(os.path.basename(__file__))\ + .getOrCreate() + +############################################################################### +## prepare data + +# read the data into a spark data frame. +cwd = os.getcwd() +filename = "AdultCensusIncome.csv" + +## NOTE: reading text file from local file path seems flaky! +#import urllib.request +#url = "https://amldockerdatasets.azureedge.net/" + filename +#local_filename, headers = urllib.request.urlretrieve(url, filename) +#datafile = "file://" + os.path.join(cwd, filename) + +data_all = spark.read.format('csv')\ + .options( + header='true', + inferSchema='true', + ignoreLeadingWhiteSpace='true', + ignoreTrailingWhiteSpace='true')\ + .load(filename) #.load(datafile) for local file + +print("Number of rows: {}, Number of coulumns : {}".format(data_all.count(), len(data_all.columns))) + +#replace "-" with "_" in column names +columns_new = [col.replace("-", "_") for col in data_all.columns] +data_all = data_all.toDF(*columns_new) + +data_all.printSchema() +data_all.show(5) + +# choose feature columns and the label column for training. +label = "income" +#xvars = ["age", "hours_per_week"] #all numeric +xvars = ["age", "hours_per_week", "education"] #numeric + string + +print("label: {}, features: {}".format(label, xvars)) + +select_cols = xvars +select_cols.append(label) +data = data_all.select(select_cols) + +############################################################################### +## split data into train and test. + +train, test = data.randomSplit([0.75, 0.25], seed=123) + +print("train ({}, {})".format(train.count(), len(train.columns))) +print("test ({}, {})".format(test.count(), len(test.columns))) + +train_data_path = os.path.join(hdfs_path, "AdultCensusIncomeTrain") +test_data_path = os.path.join(hdfs_path, "AdultCensusIncomeTest") + +# write the train and test data sets to intermediate storage and then read +train.write.mode('overwrite').orc(train_data_path) +test.write.mode('overwrite').orc(test_data_path) + +print("train and test datasets saved to {} and {}".format(train_data_path, test_data_path)) + +train_read = spark.read.orc(train_data_path) +test_read = spark.read.orc(test_data_path) + +assert train_read.schema == train.schema and train_read.count() == train.count() +assert test_read.schema == test.schema and test_read.count() == test.count() + +############################################################################### +## train model + +from pyspark.ml import Pipeline, PipelineModel +from pyspark.ml.feature import OneHotEncoderEstimator, StringIndexer, IndexToString, VectorAssembler +from pyspark.ml.classification import LogisticRegression + +# create a new Logistic Regression model, which by default uses "features" and "label" columns for training. +reg = 0.1 +lr = LogisticRegression(regParam=reg) + +# encode string columns +dtypes = dict(train.dtypes) +dtypes.pop(label) + +si_xvars = [] +ohe_xvars = [] +featureCols = [] +for idx,key in enumerate(dtypes): + if dtypes[key] == "string": + featureCol = "-".join([key, "encoded"]) + featureCols.append(featureCol) + + tmpCol = "-".join([key, "tmp"]) + si_xvars.append(StringIndexer(inputCol=key, outputCol=tmpCol, handleInvalid="skip")) #, handleInvalid="keep" + ohe_xvars.append(OneHotEncoderEstimator(inputCols=[tmpCol], outputCols=[featureCol])) + else: + featureCols.append(key) + +# string-index the label column into a column named "label" +si_label = StringIndexer(inputCol=label, outputCol='label') +#si_label._resetUid("si_label") # try to name the transformer, which seems not carried over to the fitted pipeline. + +# assemble the encoded feature columns in to a column named "features" +assembler = VectorAssembler(inputCols=featureCols, outputCol="features") + +# put together the pipeline +stages = [] +stages.extend(si_xvars) +stages.extend(ohe_xvars) +stages.append(si_label) +stages.append(assembler) +stages.append(lr) + +pipe = Pipeline(stages=stages) +print("Pipeline Created") + +# train the model +model = pipe.fit(train) +print("Model Trained") +print("Model is ", model) +print("Model Stages", model.stages) + +# name the string-index stage for the label so it can be identified easier later +model.stages[2]._resetUid("si_label") + +############################################################################### +## evaluate model + +from pyspark.ml.evaluation import BinaryClassificationEvaluator + +# make prediction +pred = model.transform(test) + +# evaluate. note only 2 metrics are supported out of the box by Spark ML. +bce = BinaryClassificationEvaluator(rawPredictionCol='rawPrediction') +au_roc = bce.setMetricName('areaUnderROC').evaluate(pred) +au_prc = bce.setMetricName('areaUnderPR').evaluate(pred) + +print("Area under ROC: {}".format(au_roc)) +print("Area Under PR: {}".format(au_prc)) + +############################################################################### +## save and load the model with ML persistence +# https://spark.apache.org/docs/latest/ml-pipeline.html#ml-persistence-saving-and-loading-pipelines + +##NOTE: by default the model is saved to and loaded from hdfs +model_name = "AdultCensus.mml" +model_fs = os.path.join(hdfs_path, model_name) + +model.write().overwrite().save(model_fs) +print("saved model to {}".format(model_fs)) + +# load the model file (from hdfs) +print("load pyspark model from hdfs") +model_loaded = PipelineModel.load(model_fs) +assert str(model_loaded) == str(model) + +print("loaded model from {}".format(model_fs)) +print("Model is " , model_loaded) +print("Model stages", model_loaded.stages) + +############################################################################### +## export and import model with mleap + +import mleap.pyspark +from mleap.pyspark.spark_support import SimpleSparkSerializer + +# serialize the model to a local zip file in JSON format +#model_name_export = "adult_census_pipeline.zip" +model_name_path = cwd +model_file = os.path.join(model_name_path, model_name_export) + +# remove an old model file, if needed. +if os.path.isfile(model_file): + os.remove(model_file) + +model_file_path = "jar:file:{}".format(model_file) +model.serializeToBundle(model_file_path, model.transform(train)) + +## import mleap model +model_deserialized = PipelineModel.deserializeFromBundle(model_file_path) +assert str(model_deserialized) == str(model) + +print("The deserialized model is ", model_deserialized) +print("The deserialized model stages are", model_deserialized.stages) + +############################################################################## +## export the final model with mleap + +## remove the stringIndexer for the label column so it won't be required for prediction +model_final = model.copy() + +si_label_index = -3 +model_final.stages.pop(si_label_index) #si_label + +## append an IndexToString transformer to the model pipeline to get the original labels +#labelReverse = IndexToString(inputCol = "label", outputCol = "predIncome") #no need to provide labels +labelReverse = IndexToString( + inputCol = "prediction", + outputCol = "predictedIncome", + labels = model.stages[si_label_index].labels) #must provide labels (from si_label) otherwise will fail +model_final.stages.append(labelReverse) + +pred_final = model_final.transform(test) +pred_final.printSchema() +pred_final.show(5) + +# remove an old model file, if needed. +if os.path.isfile(model_file): + os.remove(model_file) +model_final.serializeToBundle(model_file_path, model_final.transform(train)) + +print("persist the mleap bundle from local to hdfs") +from subprocess import Popen, PIPE +hdfs_fs_put = ["hadoop", "fs", "-put", "-f", model_file, os.path.join(hdfs_path, model_name_export)] +proc = Popen(hdfs_fs_put, stdout=PIPE, stderr=PIPE) +s_output, s_err = proc.communicate() +if (s_err): + print("s_output: {s_output}\ns_err: {s_err}".format(s_output=s_output, s_err=s_err)) + +############################################################################### diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_sql_tests.py b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_sql_tests.py new file mode 100644 index 00000000..a1b3e1c8 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/mleap_sql_tests.py @@ -0,0 +1,164 @@ +import os +dir_path = os.path.dirname(os.path.realpath(__file__)) + +import sys +sys.path.append(os.path.join(dir_path, os.pardir, os.pardir, os.pardir)) + +from spark_submit import * + +from subprocess import run, PIPE +import pytest +import pyodbc + + +@pytest.fixture(scope="module") +def setup_mod(): + print("setting up module ...") + + odbcDriver = "ODBC Driver 13 for SQL Server" + databaseName = "tempdb" + headNode = "master-0.master-svc" + + # Read sql username and password from environment variable. + username = os.environ["EXTENSIBILITY_TEST_SQL_USER"] + password = os.environ["EXTENSIBILITY_TEST_SQL_PASSWORD"] + if not username or not password: + raise Exception("Environment variable EXTENSIBILITY_TEST_SQL_USER or EXTENSIBILITY_TEST_SQL_PASSWORD cannot not be found") + + # enable SPEES + conn = pyodbc.connect('DRIVER={0};SERVER={1};DATABASE={2};UID={3};PWD={4}'.format( + odbcDriver, headNode, databaseName, username, password), autocommit=True) + cursor = conn.cursor() + cursor.execute("""EXEC sp_configure 'external scripts enabled', 1""") + assert(-1 == cursor.rowcount) + + cursor.execute("""RECONFIGURE""") + assert(-1 == cursor.rowcount) + + yield dict(cursor=cursor) + + print("tearing down module ...") + + +def test_java_spees(setup_mod): + # exectue a Java SPEES query to create external libraries + cursor = setup_mod['cursor'] + cursor.execute(""" + --SELECT @@SERVERNAME AS 'Server Name', @@VERSION AS 'Server Version', @@SERVICENAME AS 'Service Name' + + IF NOT EXISTS (SELECT * FROM sys.external_languages WHERE language = 'Java') + --DROP EXTERNAL LANGUAGE Java; + CREATE EXTERNAL LANGUAGE Java + FROM (CONTENT = N'/opt/mssql/lib/extensibility/java-lang-extension.tar.gz', file_name = 'javaextension.so'); + + IF EXISTS (SELECT * FROM sys.external_libraries WHERE name = 'SdkPackage') + DROP EXTERNAL LIBRARY SdkPackage; + CREATE EXTERNAL LIBRARY SdkPackage + FROM (CONTENT = '/opt/mssql/lib/mssql-java-lang-extension.jar') WITH (LANGUAGE = 'Java'); + + IF EXISTS (SELECT * FROM sys.external_libraries WHERE name = 'TestPackage') + DROP EXTERNAL LIBRARY TestPackage + CREATE EXTERNAL LIBRARY TestPackage + FROM (CONTENT = '/opt/mssql/java/jars/JavaTestPackage.jar') WITH (LANGUAGE = 'Java'); + + DECLARE @script NVARCHAR(max) = N'JavaTestPackage.PassThrough' --no space allowed in the string! + EXEC sp_execute_external_script + @language = N'Java' + , @script = @script + , @input_data_1 = N'SELECT 1' + """) + + rows = cursor.fetchall() + assert(1 == len(rows)) + assert(1 == rows[0][0]) + + +def dictfetchall(cursor): + '''fetch all rows from a cursor and return them as a dict''' + colnames = [col[0] for col in cursor.description] + return [dict(zip(colnames, row)) for row in cursor.fetchall()] + + +def test_mleap_pyspark(setup_mod): + # train a pyspark model and export it as a mleap bundle + hdfs_path = "/spark_ml" + model_name_export = "adult_census_pipeline.zip" + + file_path = 'mleap_pyspark.py' + file_args = [hdfs_path, model_name_export] + ret = spark_submit(file_path, file_args) + assert 0 == ret + + # get the mleap bundle from hdfs and copy it to the mssql-server container of the master-0 pod + hdfs_file_path = os.path.join(hdfs_path, model_name_export) + ret = run(["hdfs", "dfs", "-get", "-f", hdfs_file_path], stdout=PIPE, stderr=PIPE).returncode + assert 0 == ret + + local_file_path = os.path.join("master-0:", "tmp") + ret = run(["kubectl", "cp", model_name_export, local_file_path, "-c", "mssql-server"], stdout=PIPE, stderr=PIPE).returncode + assert 0 == ret + + # exectue a Java SPEES query to serve the mleap bundle + cursor = setup_mod['cursor'] + cursor.execute(""" + --suppresses the record count values generated by DML statements + --like UPDATE and allows the result set to be retrieved directly. + SET NOCOUNT ON; + + IF EXISTS (SELECT * FROM sys.external_libraries WHERE name = 'MleapApp') + DROP EXTERNAL LIBRARY MleapApp; + CREATE EXTERNAL LIBRARY MleapApp + FROM (CONTENT = '/opt/mssql/java/jars/mssql-mleap-app-assembly-1.0.jar') WITH (LANGUAGE = 'Java') + + DROP TABLE IF EXISTS ##test + CREATE TABLE ##test ( + income nvarchar(10) + , age int + , hours_per_week int + , education nvarchar(10) + , sex nvarchar(10) + ); + INSERT INTO ##test values ('<=50K', 39, 40, 'Bachelors', 'Male'); + INSERT INTO ##test values ('<=50K', 50, 13, 'Bachelors', 'Male'); + INSERT INTO ##test values ('<=50K', 38, 40, 'HS-grad', 'Male'); + --SELECT * FROM ##test + + DECLARE @script NVARCHAR(max) = N'com.microsoft.sqlserver.mleap.Scorer' --no space allowed in the string! + DECLARE @language nvarchar(4) = N'Java' + DECLARE @parallel bit = 0 + DECLARE @input_data_1 nvarchar(97) = N'select age, hours_per_week, education, sex, income from ##test' + DECLARE @params nvarchar(200) = N'@modelPath nvarchar(100), @outputFields nvarchar(100), @logLevel nvarchar(100)' + DECLARE @modelPath nvarchar(100) = N'/tmp/adult_census_pipeline.zip' + DECLARE @outputFields nvarchar(100) = N'prediction,probability,education,sex,income,predictedIncome' + DECLARE @logLevel nvarchar(100) = N'INFO' + EXEC sp_execute_external_script @language = @language, @script = @script, @parallel = @parallel + , @input_data_1 = @input_data_1 + , @params = @params, @modelPath = @modelPath, @outputFields = @outputFields, @logLevel = @logLevel + WITH RESULT SETS ((prediction int, probability0 float, probability1 float, education nvarchar(20), sex nvarchar(20), income nvarchar(20), predictedIncome nvarchar(20))) + """) + + rows = dictfetchall(cursor) + #pandas.DataFrame(rows) + + assert rows == [ + {'education': 'Bachelors', + 'income': '<=50K', + 'predictedIncome': '<=50K', + 'prediction': 0, + 'probability0': 0.6544871023375456, + 'probability1': 0.3455128976624544, + 'sex': 'Male'}, + {'education': 'Bachelors', + 'income': '<=50K', + 'predictedIncome': '<=50K', + 'prediction': 0, + 'probability0': 0.7363751868447964, + 'probability1': 0.2636248131552036, + 'sex': 'Male'}, + {'education': 'HS-grad', + 'income': '<=50K', + 'predictedIncome': '<=50K', + 'prediction': 0, + 'probability0': 0.8324466132959966, + 'probability1': 0.16755338670400344, + 'sex': 'Male'}] diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/setup.sh b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/setup.sh new file mode 100644 index 00000000..80f376c3 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/setup.sh @@ -0,0 +1,17 @@ +#!/bin/bash -e + +echo "Setting up mleap_sql tests" + +export PYSPARK_PYTHON=python3 + +export EXTENSIBILITY_TEST_SQL_USER=sa +export EXTENSIBILITY_TEST_SQL_PASSWORD=Yukon900 + +hadoop fs -mkdir -p /user/root +wget https://amldockerdatasets.azureedge.net/AdultCensusIncome.csv +hadoop fs -copyFromLocal AdultCensusIncome.csv /user/root + +# Copy java ext jars to mssql-server container in master pod +#kubectl cp -c mssql-server ../jars/mssql_java_lang_extension.jar master-0:/opt/mssql/java/jars/ +kubectl cp -c mssql-server ../jars/JavaTestPackage.jar master-0:/opt/mssql/java/jars/ +kubectl cp -c mssql-server ../mssql-mleap-app/target/scala-2.11/mssql-mleap-app-assembly-1.0.jar master-0:/opt/mssql/java/jars/ diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/test.sh b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/test.sh new file mode 100644 index 00000000..7308e25f --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mleap_sql_test/test.sh @@ -0,0 +1,6 @@ +#!/bin/bash + +source ./setup.sh + +# Generate Junit results +python3 -m pytest -v --junitxml /tests/junit/mleap_sql.xml -o junit_suite_name=mleap_sql --durations=0 mleap_sql_tests.py diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/Makefile b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/Makefile new file mode 100644 index 00000000..5f7abc4e --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/Makefile @@ -0,0 +1,16 @@ +.DEFAULT_GOAL = all + +all: assembly + +assembly: + @sbt assembly + +package: + @sbt package + +clean: + @rm -rf project/project + @rm -rf project/target + @rm -rf target + @rm -rf .idea + diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/build.sbt b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/build.sbt new file mode 100644 index 00000000..ef27c126 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/build.sbt @@ -0,0 +1,23 @@ +name := "mssql-mleap-app" + +version := "1.0" + +scalaVersion := "2.11.12" + +libraryDependencies ++= Seq( + "ml.combust.mleap" %% "mleap-runtime" % "0.13.0" % "provided", + "org.apache.commons" % "commons-csv" % "1.5", + "commons-cli" % "commons-cli" % "1.4", + "org.scalatest" %% "scalatest" % "3.2.0-SNAP10" % Test, + "org.scalacheck" %% "scalacheck" % "1.14.0" % Test, + "com.novocode" % "junit-interface" % "0.11" % Test +) + +// Exclude scala-library from this fat jar. The scala library is already there in spark package. +assemblyOption in assembly := (assemblyOption in assembly).value.copy(includeScala = false) + +// exclude specific jars +assemblyExcludedJars in assembly := { + val cp = (fullClasspath in assembly).value + cp filter {_.data.getName == "mssql_java_lang_extension.jar"} +} diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/lib/mssql_java_lang_extension.jar b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/lib/mssql_java_lang_extension.jar new file mode 100644 index 0000000000000000000000000000000000000000..85cb713358b76adf4aed4542ceb5329840bd4508 GIT binary patch literal 4748 zcmb_gc{r5q9v*`%*~xHHwn!PIv4w0I+Zg*YCCix5V8|{@_GPS5_Ut4vWLF4TDoa9T z8EdwXWlUrchx46tRpLVGbt;BLf_c$th?7 z0BXRWrc7~|fpW%({W^o;-x(M*%GKG;+1C9xxg3AWb#rlavvKvbas5rz`QKF8BR!Fh zNGCfRZ+9CfH?*_Uzi|F~`IY`2jw>32c1L^Ks3P5wZZ__sC`Y858!nmT%L?V{oP@{I z^IjIZ4;0~0s=%q#>(dHhibKspcyLts8=4@gh80_K94)QT6qf4rR8Z6c@lm|D9#gqG zcK{)NgOEASm4P`w=Zz4onw)_^*1T+M5B7G6&Vb~Z5~ge7XWD3a19LU&SvmWM^Ik{I z=iFq#+RB7p0dyg)2P@ygp}b z9wZWi0Zl?P0-R#q-q6mHi#pxj@XqCn6L$#@V<5_n0y7xSHI38<7d9P7haVQaj|qJy zvYQm5CE?)U<6d%cm=|eq9mW}zSz|;a3~O>2Gc^}C=ubtKoO{2tp8B!2kh3_><5ZZk zXB=IWvWYoJ(3$0dPfAsr0qZq22>pu}W#_W=+*^~!hk7=3Jbe%DNFwDz-p#YUUXmqh zT)Txauq=O7YLU;G-lU@fb-;_#B$?Y9oJhA1eRc+Q;d zo=OF+2rY$0Emcr-C|?gcVDkfDHOS!#N-&e*|Fo(0bE*Klz^KX)0Ni^ZGL#$ z&>>-9F(f!-3^R>)Ixq0{^pL5#WnEe!pM{;5xVVDVH32NaI!{<_VnlTTQMgf&nVyT_ zCn3@*{4+_ny@WlH+h|HHTa9>HUG+NzaI-xc>vzMA_abhXykAG;eXmNviE&M8L+kjs zxW=a&J~pnOGsY5=ZqCa1Tvx{)i{oGY@YXa=HeIX3J-ESKxA@v%SkYJ~OKzdHarE2g zI$j$VxQFbVloexGQbD>g-4_R_mbDU`Yn@x5BIz7T$jKHLI+~BgjMcXd45Z^u8C|-c zsJ{XQa(2>|&fW%Rt;K5D24={{gxzG!hfC=9gV}po`oXNd4Kl(r*E!av{b`8uUxUg@ z&Ot4m3+rjyCN4bCqrxosNs=~~dhC4+YU*F?_3-a%3xBjwNI-bq*k97%;GVkD|DdRz z8Cqu~aEV1fq9!JjW`^Gn%)YM40a{0MWFs*gXbz-boY>@F-SArUeG$Vw&55}39GGa`UjTeH1V$^ z!HN2~kjE!ZP>5&AI+ldj=jQDLwT9&hvzy;I@dhlOej4#|;wLxCfN)zKlkw+$}_u8NhY$e*Cszjb*8%R2DjDRZ)lEQ&B4qc zm&4V3hE5`!l(J%?TEQSc-8542)m8+nE<(*WG2$_wK&Y9X)A(Fva5{30g%2fqa05MO zsqucV4|0sgs{+(A?g;Xk09C~9On`PLQoJWpD$VC_TgHuiCdO2mL%R(hSuuL&__G=8 z8dGe=tjTrmn=O2-xN@KuDlr@K}(oPA%oB2da$o z7d9+P-@s@2_&YvE-~&``h&G{RR;~a$@q_a(e9Q8}$=^x6Ggt8^3TUHmS}m|VjI$nH zBW;i9=%6S)2Mq+kULcE-4hK}Bw2G&;G36js{j@mo(%q4XI3MtoBmS)>>2!ggCi@Ii zbic9~)T{?v-a0z=Wr1M&iGyYEY49S;b5RZb6>rr&*|C&p%dS@7UgpGveAL6em~tyM zK!DY}SJ>EH)*3G3(zMXs3Wbzrz7$1Bw&lv}Y!8}V%U!0gEWFjLT->~tJ+3~#IG z@y=wMfMVHVv;n=3=lb55yho%_^gFs@smqKNcSsmSYk(5!l0CEGGQ}xA+tgI%YhIOK zqyxRvBimHDay>wvlhhxDm|sT9=q&S(U6c-#QEu1d)9D}W9+!HPKHjC!|)(>e_gpWj|2BUVNi7BKBZvQ&ko zEzbd&vj-BVY!u1buPfYXPf_)QMt8YMlC{epj-*xnV6d)Bi9@PA3p46|y6ueBH-{CO z+gA=@^!Y>wT?5r`10&@QTQoBfgniKq>`ajsSvRJy8f+vfO209^qf8LC+>>v}IMBV5 zb?{Mo(|K)S*gg2z0r`;}U6CkvLl;NG zfAKr?KLic+wn2HgJG=hDr%i@%=RqZ=Ot|(~=ojJK>T}f|$0N^BfzLP@^z+TtmZi}Zi@G&=L8VcO zc{5&6qU)lbBJwJmAjOAa8aGzX&JWld*S&eY5T#r}v9F7y;#o*BTIW+8Nk);P^c+D3 z3wKW&4i@>{7`?c){Ycr$O=U^I^}KN1vWxC2^D7ev1~S|U2Jwbrg3PjDYhVYJ`xHve z@B*;p{j%{WrsG|oFz=V@QJ1zBI(caZ=(ckD$5$Y*f^9`{eN>m=U|Mg`U2%^>|5^OG zdFyqT@>pUp_KP2D!%lm+Sn>9Dx?1e0b~n>2{6wT-TDweJ-vOj;brZ9nut?94CVbKt zl5OPWuXdY-S!fu^((w6d(LSNA*$qeFylg-j{e+?$ni?YSDvI~-hO$7B;E^Ywiqx+A z3zZrVzwUmSTG4I~~TEf%i zZ7oTJNj|wKV?NzjyEZdw)8;h0L9@)3DuxzAlP_ycv!n8jPQzR-tjBJ1J~zgb?A@+d zd{S3x6QwZ9LcmyMRzN<57MIr)D|PRMr=ALy?YYX9zouNs$;tAmA*2YVKDd7eYL>(% z3B7ZCCG-BRuhK_bwn|Lt2?PiLkPh|fku?45Evx^3+Oj_>z}Tb}FIFg9XZDI&=de^E z<$0dO+Qx#MSiU@9M#ODiWL=(llp2Qc%6%I=^&_ft+}GwW7%01n+uEmaElP!sDsn=E#ob*U8#X#Ontz z+1e(4vtxp2hq+fR5=lteB7Jf`JCmNS1ch}jC~cX8@SgnxTB$r|EgZJ`G)F>z%Q3JXGLD(7%nyjJppRI{zY0iZUG z#%-ePi(ahSOpPwO+JKsoTf{#%JKwY+$bDqRThh4l3y*mK0<)V{N}6MBC~X|-!u!n# ztLt)4SK{_0n*#p9JslNv@*9y$W<{q2>QS;|0lPg=bSlHy{_z}Dh_A;h`Ku+F9@iBe zB+Y8=&qqitl}Z%DTxKrh1UdFV!LFELC2$>^G;4P$-43>}OcmF$z*)^x2sFmr8!@d?87vO`LAqwTFo~~m2Q^&)B|AuXWF_US@)zTTema!bGPp4P)4H%Y)T*Gb#oMXT>Hj~iXhIcT_Ah~ zWb-yfIWSLrBHW{ER`NL|&9{e$&$ix%4se-VRp?kx(Qp=A51#fgbP`x!6K`DYdAU6s zKzueJIm^GN8m|l4`^RF@cFf$Q-8R0Uj&P~_@QwL^+;;Gs9moFZMRUK`8&!18Iy)zx z+mh(oKpJb)cS5<>2yS$ieExyXh<&;?=Y+Buv*c{MZS#S(^}9`zl;MgzB32=DeHZt5 z(aag(fVk=Ggyi+h-Pv~Dnmbos_=#R@uG0I>m2PN5{yALh4Vi$N_+Sx}#%+49Euf?i_#g$4AEf=lGS4KSH4e z{Mt)@N2&Vv;q<$TBX9jZ-XA_5dGQ~e_OH7gdF}5(cIfhduKm$@{~G1Udw&n<-=O@J z6aRTdf7BUD-$4(b0{$NEUpn(2|KE{z`H>=yw9EHsI=tuj)mr@dEB;p{^WXh{+*?OP h^gV7J|Nr#=iI((Wl!p@T=njV((0s^@b&N-&e*y?;?Gyk2 literal 0 HcmV?d00001 diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/build.properties b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/build.properties new file mode 100644 index 00000000..e9a676f7 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/build.properties @@ -0,0 +1 @@ +sbt.version = 1.1.5 \ No newline at end of file diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/plugins.sbt b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/plugins.sbt new file mode 100644 index 00000000..652a3b93 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/project/plugins.sbt @@ -0,0 +1 @@ +addSbtPlugin("com.eed3si9n" % "sbt-assembly" % "0.14.6") diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/PrimitiveDataset.java b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/PrimitiveDataset.java new file mode 100644 index 00000000..3e9a269d --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/PrimitiveDataset.java @@ -0,0 +1,85 @@ +package com.microsoft.sqlserver.mleap; + +import java.sql.JDBCType; +import java.sql.Types; +import java.util.Arrays; + +public class PrimitiveDataset extends com.microsoft.sqlserver.javalangextension.PrimitiveDataset { + public String[] getColumnNames() { + int nCols = getColumnCount(); + String[] columnNames = new String[nCols]; + + for (int iCol = 0; iCol < nCols; iCol++) { + columnNames[iCol] = getColumnName(iCol); + } + + return columnNames; + } + + public int[] getColumnTypes() { + int nCols = getColumnCount(); + int[] columnTypes = new int[nCols]; + + for (int iCol = 0; iCol < nCols; iCol++) { + columnTypes[iCol] = getColumnType(iCol); + } + + return columnTypes; + } + + public int getColumnIndex(String columnName) { + String[] columnNames = getColumnNames(); + int index = Arrays.asList(columnNames).indexOf(columnName); + return index; + } + + public int getRowCount(int iCol) { + int sqlType = getColumnType(iCol); + int columnLength; + + switch(sqlType) { + case Types.BIT: + columnLength = getBooleanColumn(iCol).length; + break; + case Types.SMALLINT: + columnLength = getShortColumn(iCol).length; + break; + case Types.INTEGER: + columnLength = getIntColumn(iCol).length; + break; + case Types.BIGINT: + columnLength = getLongColumn(iCol).length; + break; + case Types.FLOAT: + columnLength = getFloatColumn(iCol).length; + break; + case Types.DOUBLE: + columnLength = getDoubleColumn(iCol).length; + break; + case Types.NVARCHAR: + columnLength = getStringColumn(iCol).length; + break; + case Types.VARBINARY: + columnLength = getBinaryColumn(iCol).length; + break; + case Types.DATE: + columnLength = getDateColumn(iCol).length; + break; + default: + throw new IllegalArgumentException("unsupported sql type: " + JDBCType.valueOf(sqlType).getName()); + } + + return columnLength; + } + + public int[] getRowCounts() { + int nCols = getColumnCount(); + int[] rowCounts = new int[nCols]; + + for (int iCol = 0; iCol < nCols; iCol++) { + rowCounts[iCol] = getRowCount(iCol); + } + + return rowCounts; + } +} diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/Scorer.java b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/Scorer.java new file mode 100644 index 00000000..a7f60117 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/java/com/microsoft/sqlserver/mleap/Scorer.java @@ -0,0 +1,314 @@ +package com.microsoft.sqlserver.mleap; + +import com.microsoft.sqlserver.javalangextension.AbstractSqlServerExtensionExecutor; + +import org.apache.commons.csv.CSVFormat; +import org.apache.commons.csv.CSVRecord; +import org.apache.commons.cli.*; + +import java.io.BufferedReader; +import java.io.FileReader; +import java.io.Reader; + +import java.sql.JDBCType; +import java.sql.Types; + +import java.util.Arrays; +import java.util.List; +import java.util.LinkedHashMap; +import java.util.logging.Level; +import java.util.logging.Logger; + +public class Scorer extends AbstractSqlServerExtensionExecutor { + + private static final Logger LOGGER = Logger.getLogger(Scorer.class.getName()); + + public Scorer() { + executorExtensionVersion = SQLSERVER_JAVA_LANG_EXTENSION_V1; + executorInputDatasetClassName = PrimitiveDataset.class.getName(); + executorOutputDatasetClassName = PrimitiveDataset.class.getName(); + } + + public void init(String sessionId, int taskId, int numTasks) { + System.out.println("init SessionID: " + sessionId + " taskId: " + taskId + " numTasks: " + numTasks); + } + + public PrimitiveDataset execute(PrimitiveDataset input, LinkedHashMap params) { + List logLevels = Arrays.asList("OFF", "SEVERE", "WARNING", "INFO", "CONFIG", "FINE", "FINER", "FINEST", "ALL"); + String logLevel = params.getOrDefault("logLevel", "WARNING").toString(); + if (!logLevels.contains(logLevel)) { + throw new IllegalArgumentException("logLevel (" + logLevel + ") must be one of " + logLevels.toString()); + } + LOGGER.setLevel(Level.parse(logLevel)); + + LOGGER.info("Logger Name: " + LOGGER.getName() + "; Logger Level:" + LOGGER.getLevel()); + + // load model + String modelPath; + try { + modelPath = params.get("modelPath").toString(); + } catch (NullPointerException e) { + throw new IllegalArgumentException("modelPath parameter is required but not set."); + } + + long startTime = System.nanoTime(); + Predictor scorer = new Predictor(); + scorer.init(modelPath); + long endTime = System.nanoTime(); + long duration = (endTime - startTime); //divide by 10^6 to get milliseconds. + LOGGER.info("model loading time: " + duration/1e6 + " ms"); + + // convert PrimitiveDataset to DefaultLeapFrame + startTime = System.nanoTime(); + scorer.primitiveDataset2leapFrame(input); + endTime = System.nanoTime(); + duration = (endTime - startTime); //divide by 10^6 to get milliseconds. + LOGGER.info("PrimitiveDataset to DefaultLeapFrame conversion time: " + duration/1e6 + " ms"); + + // do prediction + startTime = System.nanoTime(); + scorer.run(); + endTime = System.nanoTime(); + duration = (endTime - startTime); //divide by 10^6 to get milliseconds. + LOGGER.info("model scoring time: " + duration/1e6 + " ms"); + + //select output fields specified + startTime = System.nanoTime(); + String[] outputFields; + outputFields = params.getOrDefault("outputFields", "").toString().split(","); + scorer.select(outputFields); + endTime = System.nanoTime(); + duration = (endTime - startTime); //divide by 10^6 to get milliseconds. + LOGGER.info("data selection time: " + duration/1e6 + " ms"); + + // convert DefaultLeapFrame to PrimitiveDataset + startTime = System.nanoTime(); + PrimitiveDataset output = new PrimitiveDataset(); + scorer.leapFrame2primitiveDataset(output); + endTime = System.nanoTime(); + duration = (endTime - startTime); //divide by 10^6 to get milliseconds. + LOGGER.info("DefaultLeapFrame to PrimitiveDataset conversion time: " + duration/1e6 + " ms"); + + return output; + } + + public void cleanup() { + System.out.println("\n* cleanup"); + } + + /** + * + * @param args commandline options for model path and input file. + * @throws Exception + *
+     * {@code
+     *
+     * -- ex: Linux
+     * java -cp mssql-mleap-app-assembly-1.0.jar:mssql_java_lang_extension.jar:mssql-mleap-lib-assembly-1.0.jar:commons-csv-1.5.jar:commons-cli-1.4.jar com.microsoft.sqlserver.mleap.Scorer
+     *  -m /tmp/adult_census_pipeline.zip
+     *  -i /tmp/adult_census_income.csv
+     *
+     * java -cp "*" -m /tmp/adult_census_pipeline.zip -i /tmp/adult_census_income.csv
+     *
+     * -- ex: Windows
+     * java -cp mssql-mleap-app-assembly-1.0.jar;mssql_java_lang_extension.jar;mssql-mleap-lib-assembly-1.0.jar;commons-csv-1.5.jar:commons-cli-1.4.jar com.microsoft.sqlserver.mleap.Scorer
+     *  -m C:\\Users\\lgong\\Work\\git\\aml-databricks\\examples\\mleapsql2\\src\\main\\resources\\sqlqueries\\adult_census_pipeline.zip
+     *  -i C:\\Users\\lgong\\Work\\git\\aml-databricks\\examples\\mleapsql2\\src\\main\\resources\\sqlqueries\\adult_census_income.csv
+     *
+     * java -cp "*"
+     *  -m C:\\Users\\lgong\\Work\\git\\aml-databricks\\examples\\mleapsql2\\src\\main\\resources\\sqlqueries\\adult_census_pipeline.zip
+     *  -i C:\\Users\\lgong\\Work\\git\\aml-databricks\\examples\\mleapsql2\\src\\main\\resources\\sqlqueries\\adult_census_income.csv
+     * }
+     * 
+ */ + public static void main(String[] args) throws Exception { + // get model and testing data + Options options = new Options(); + + Option input = new Option("i", "input", true, "input file"); + input.setRequired(true); + options.addOption(input); + + Option model = new Option("m", "model", true, "model path"); + model.setRequired(true); + options.addOption(model); + + CommandLineParser parser = new DefaultParser(); + HelpFormatter formatter = new HelpFormatter(); + CommandLine cmd = null; + + try { + cmd = parser.parse(options, args); + } catch (ParseException e) { + System.out.println(e.getMessage()); + formatter.printHelp("Scorer", options); + + System.exit(1); + } + + String modelPath = cmd.getOptionValue("model"); + String scoreFile = cmd.getOptionValue("input"); + + LOGGER.info("os.name: " + System.getProperty("os.name")); + LOGGER.info("isWindows: " + System.getProperty("os.name").startsWith("Windows")); + LOGGER.info("args: " + Arrays.toString(args)); + + LOGGER.info("modelPath: " + modelPath); + LOGGER.info("scoreFile: " + scoreFile); + + // read in the testing data + BufferedReader bufferedReader = new BufferedReader(new FileReader(scoreFile)); + int nRows = -1; //account for the header row + while(bufferedReader.readLine() != null) { + nRows++; + } + + LinkedHashMap inputFields = new LinkedHashMap(); + inputFields.put("age", Types.INTEGER); + inputFields.put("workclass", Types.NVARCHAR); + inputFields.put("fnlwgt", Types.INTEGER); + inputFields.put("education", Types.NVARCHAR); + inputFields.put("education_num", Types.INTEGER); + inputFields.put("marital_status", Types.NVARCHAR); + inputFields.put("occupation", Types.NVARCHAR); + inputFields.put("relationship", Types.NVARCHAR); + inputFields.put("race", Types.NVARCHAR); + inputFields.put("sex", Types.NVARCHAR); + inputFields.put("capital_gain", Types.INTEGER); + inputFields.put("capital_loss", Types.INTEGER); + inputFields.put("hours_per_week", Types.INTEGER); + inputFields.put("native_country", Types.NVARCHAR); + inputFields.put("income", Types.NVARCHAR); + + String[] columnNames = {"age", "hours_per_week", "education", "sex", "income"}; //choose the input variables + int[] columnTypes = new int[columnNames.length]; + for (int iCol = 0; iCol < columnNames.length; iCol++) { + try { + columnTypes[iCol] = inputFields.get(columnNames[iCol]); + } catch (NullPointerException e) { + throw new IllegalArgumentException("invalid input field: " + columnNames[iCol]); + } + } + int nCols = columnNames.length; + + Object[] columns = new Object[nCols]; + for (int iCol = 0; iCol < nCols; iCol++) { + int columnType = columnTypes[iCol]; + switch (columnType) { + case Types.INTEGER: + columns[iCol] = new int[nRows]; + break; + case Types.NVARCHAR: + columns[iCol] = new String[nRows]; + break; + default: + throw new IllegalArgumentException("unsupported sql type: " + JDBCType.valueOf(columnType).getName()); + } + } + + Reader in = new FileReader(scoreFile); + Iterable records = CSVFormat.RFC4180.withFirstRecordAsHeader().parse(in); + int iRow = 0; + for (CSVRecord record : records) { + for (int iCol = 0; iCol < nCols; iCol++) { + int columnType = columnTypes[iCol]; + switch (columnType) { + case Types.INTEGER: + ((int[])(columns[iCol]))[iRow] = Integer.parseInt(record.get(columnNames[iCol])); + break; + case Types.NVARCHAR: + ((String[])(columns[iCol]))[iRow] = record.get(columnNames[iCol]); + break; + default: + throw new IllegalArgumentException("unsupported sql type: " + JDBCType.valueOf(columnType).getName()); + } + } + iRow++; + } + + // form the primitive dataset + PrimitiveDataset inputds = new PrimitiveDataset(); + + for (int iCol = 0; iCol < nCols; iCol++) { + int columnType = columnTypes[iCol]; + switch (columnType) { + case Types.INTEGER: + inputds.addColumnMetadata(iCol, columnNames[iCol], Types.INTEGER, 0, 0); + inputds.addIntColumn(iCol, (int[])(columns[iCol]), null); + break; + case Types.NVARCHAR: + inputds.addColumnMetadata(iCol, columnNames[iCol], Types.NVARCHAR, 0, 0); + inputds.addStringColumn(iCol, (String[])(columns[iCol])); + break; + default: + throw new IllegalArgumentException("unsupported sql type: " + JDBCType.valueOf(columnType).getName()); + } + } + + // specify some params + LinkedHashMap params = new LinkedHashMap<>(); + params.put("logLevel", "INFO"); //default WARN + params.put("modelPath", modelPath); + params.put("outputFields", "prediction,probability,education,sex,income,predictedIncome"); + //params.put("outputFields", "features,education-encoded"); //SparseTensor + + // perform scoring + Scorer scorer = new Scorer(); + scorer.init("session0", 0, 1); + PrimitiveDataset output = scorer.execute(inputds, params); + + // display output + int nOutputCols = output.getColumnCount(); + System.out.println("\nnOutputCols: " + nOutputCols); + + for (int iCol = 0; iCol < nOutputCols; iCol++) { + System.out.println("\nColumnName: " + output.getColumnName(iCol)); + + int columnType = output.getColumnType(iCol); + System.out.println("ColumnType: " + JDBCType.valueOf(columnType).getName()); + + switch(columnType) { + case Types.INTEGER: + int[] intColumn = output.getIntColumn(iCol); + System.out.println("Column.length: " + intColumn.length); + System.out.println("Column: " + Arrays.toString(intColumn)); + break; + case Types.DOUBLE: + double[] doubleColumn = output.getDoubleColumn(iCol); + System.out.println("Column.length: " + doubleColumn.length); + System.out.println("Column: " + Arrays.toString(doubleColumn)); + break; + case Types.BIGINT: + long[] longColumn = output.getLongColumn(iCol); + System.out.println("Column.length: " + longColumn.length); + System.out.println("Column: " + Arrays.toString(longColumn)); + break; + case Types.BIT: + boolean[] booleanColumn = output.getBooleanColumn(iCol); + System.out.println("Column.length: " + booleanColumn.length); + System.out.println("Column: " + Arrays.toString(booleanColumn)); + break; + case Types.FLOAT: + float[] floatColumn = output.getFloatColumn(iCol); + System.out.println("Column.length: " + floatColumn.length); + System.out.println("Column: " + Arrays.toString(floatColumn)); + break; + case Types.SMALLINT: + short[] shortColumn = output.getShortColumn(iCol); + System.out.println("Column.length: " + shortColumn.length); + System.out.println("Column: " + Arrays.toString(shortColumn)); + break; + case Types.NVARCHAR: + String[] stringColumn = output.getStringColumn(iCol); + System.out.println("Column.length: " + stringColumn.length); + System.out.println("Column: " + Arrays.toString(stringColumn)); + break; + default: + System.out.println("No columnType " + JDBCType.valueOf(columnType).getName()); + } + } + + // cleanup + scorer.cleanup(); + } +} diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/resources/adult_census_income.csv b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/resources/adult_census_income.csv new file mode 100644 index 00000000..48a1b390 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/resources/adult_census_income.csv @@ -0,0 +1,4 @@ +age,workclass,fnlwgt,education,education_num,marital_status,occupation,relationship,race,sex,capital_gain,capital_loss,hours_per_week,native_country,income +39,State-gov,77516,Bachelors,13,Never-married,Adm-clerical,Not-in-family,White,Male,2174,0,40,United-States,<=50K +50,Self-emp-not-inc,83311,Bachelors,13,Married-civ-spouse,Exec-managerial,Husband,White,Male,0,0,13,United-States,<=50K +38,Private,215646,HS-grad,9,Divorced,Handlers-cleaners,Not-in-family,White,Male,0,0,40,United-States,<=50K diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Predictor.scala b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Predictor.scala new file mode 100644 index 00000000..c319a510 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Predictor.scala @@ -0,0 +1,66 @@ +package com.microsoft.sqlserver.mleap + +import java.io.File +import java.util.logging.Logger +import java.util.logging.Level + +import ml.combust.bundle.BundleFile +import ml.combust.mleap.runtime.MleapSupport._ +import ml.combust.mleap.runtime.frame.Transformer +import resource._ + +class Predictor extends Score { + var model: Transformer = null + private val LOGGER = Logger.getLogger(classOf[Scorer].getName) + + def init(model_path: String) { + LOGGER.info(s"init($model_path)") + + model = (for(bf <- managed(BundleFile(new File(model_path)))) yield { + bf.loadMleapBundle() + }).tried.flatMap(identity).get.root + + if (LOGGER.getLevel.intValue() <= Level.INFO.intValue()) { + println("\nmodel schema fields:") + model.schema.fields.zipWithIndex.foreach { + case (field, idx) => println(s"$idx $field") + } + + println("\nmodel inputSchema fields:") + model.inputSchema.fields.zipWithIndex.foreach { + case (field, idx) => println(s"$idx $field") + } + + println("\nmodel outputSchema fields:") + model.outputSchema.fields.zipWithIndex.foreach { + case (field, idx) => println(s"$idx $field") + } + } + + LOGGER.info(s"model loaded...\n") + } + + def run(): Unit = { + frame_out = model.transform(frame_in).get + + if (LOGGER.getLevel.intValue() <= Level.INFO.intValue()) { + println("\noutput schema fields:") + frame_out.schema.fields.zipWithIndex.foreach { + case (field, idx) => println(s"$idx $field") + } + } + + //leapFrame2json(frame_out) + } + + def select(fieldNames: Array[String]) { + if (fieldNames.nonEmpty && fieldNames.length != 1 && fieldNames(0) != "") { + val allFieldNames = frame_out.schema.fields.map(_.name) + if (!fieldNames.forall(allFieldNames.contains)) { + throw new IllegalArgumentException(s"${fieldNames.toList} must be a subset of $allFieldNames") + } + + frame_out = frame_out.select(fieldNames: _*).get + } + } +} diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Score.scala b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Score.scala new file mode 100644 index 00000000..342dcfa7 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/main/scala/com/microsoft/sqlserver/mleap/Score.scala @@ -0,0 +1,235 @@ +package com.microsoft.sqlserver.mleap + +import java.io.File +import java.sql.{JDBCType, Types} + +import ml.combust.mleap.runtime.MleapSupport._ +import ml.combust.mleap.runtime.frame.{DefaultLeapFrame, Row} +import ml.combust.mleap.runtime.serialization.{BuiltinFormats, FrameReader} +import ml.combust.mleap.core.types._ +import ml.combust.mleap.tensor.{ByteString, DenseTensor, SparseTensor} + +trait Score { + + var frame_in: DefaultLeapFrame = null + var frame_out: DefaultLeapFrame = null + + def getScalaType(sqlType: Int): ScalarType = { + sqlType match { + case Types.BIT => ScalarType.Boolean + case Types.TINYINT => ScalarType.Byte + case Types.SMALLINT => ScalarType.Short + case Types.INTEGER => ScalarType.Int + case Types.BIGINT => ScalarType.Long + case Types.FLOAT => ScalarType.Float + case Types.DOUBLE => ScalarType.Double + case Types.NVARCHAR => ScalarType.String + case Types.BINARY => ScalarType.ByteString + case _ => throw new IllegalArgumentException("unsupported sql type: " + JDBCType.valueOf(sqlType).getName) + } + } + + def getSqlType(mleapType: BasicType): Int = { + mleapType match { + case BasicType.Boolean => Types.BIT + case BasicType.Byte => Types.TINYINT + case BasicType.Short => Types.SMALLINT + case BasicType.Int => Types.INTEGER + case BasicType.Long => Types.BIGINT + case BasicType.Float => Types.FLOAT + case BasicType.Double => Types.DOUBLE + case BasicType.String => Types.NVARCHAR + case BasicType.ByteString => Types.BINARY + case _ => throw new IllegalArgumentException("unsupported mleap type: " + mleapType) + } + } + + def primitiveDataset2leapFrame(input: PrimitiveDataset) { + val nCols = input.getColumnCount() + val nRows = input.getRowCount(0) // assuming columns have the same length + + // Create a schema. + val fields = List.newBuilder[StructField] + for (iCol <- 0 until nCols) { + fields += StructField(input.getColumnName(iCol), getScalaType(input.getColumnType(iCol))) + } + val schema = StructType(fields.result).get + + // Create a dataset to contain all of our values + val seqBuilder = Seq.newBuilder[Row] + for (iRow <- 0 until nRows) { + val values = List.newBuilder[Any] + for (iCol <- 0 until nCols) { + val columnType = input.getColumnType(iCol) + values += (columnType match { + case Types.BIT => input.getBooleanColumn(iCol)(iRow) + case Types.SMALLINT => input.getShortColumn(iCol)(iRow) + case Types.INTEGER => input.getIntColumn(iCol)(iRow) + case Types.BIGINT => input.getLongColumn(iCol)(iRow) + case Types.FLOAT => input.getFloatColumn(iCol)(iRow) + case Types.DOUBLE => input.getDoubleColumn(iCol)(iRow) + case Types.NVARCHAR => input.getStringColumn(iCol)(iRow) + case Types.VARBINARY => input.getBinaryColumn(iCol)(iRow) + case Types.DATE => input.getDateColumn(iCol)(iRow) + case _ => throw new IllegalArgumentException(s"No BasicType $columnType") + }) + } + seqBuilder += Row(values.result: _*) + } + + val dataset = seqBuilder.result + + // Create a LeapFrame from the schema and dataset + frame_in = DefaultLeapFrame(schema, dataset) + } + + def leapFrame2primitiveDataset(output: PrimitiveDataset) { + + val nRows = frame_out.dataset.length + var nCols = 0 + + val schema = frame_out.schema + val fields = schema.fields + + println("\nouput columns:") + for (iField <- 0 until fields.length) { + val field: StructField = fields(iField) + val name = field.name + val dataType = field.dataType + val base = dataType.base + val shape = dataType.shape + + if (shape.isTensor) { + val nDims = shape.asInstanceOf[TensorShape].dimensions.get.length + frame_out.dataset(0).getTensor(iField) match { + case dense: DenseTensor[_] => { + println(s"\t$name: DenseTensor[$base]") + } + case sparse: SparseTensor[_] => { + println(s"\t$name: SparseTensor[$base]") + } + } + + for (iDim <- 0 until nDims) { + val nSlots = field.dataType.shape.asInstanceOf[TensorShape].dimensions.get(iDim) + + for (iSlot <- 0 until nSlots) { + output.addColumnMetadata(nCols, name + iSlot, getSqlType(base), 0, 0) + base match { + case BasicType.Boolean => { + val outputDataCol = frame_out.dataset.map(_.getTensor[Boolean](iField).toDense.values(iSlot)).toArray + output.addBooleanColumn(nCols, outputDataCol, null) + } + case BasicType.Byte => { + val outputDataCol = frame_out.dataset.map(_.getTensor[Byte](iField).toDense.values(iSlot).toShort).toArray + output.addShortColumn(nCols, outputDataCol, null) + } + case BasicType.Short => { + val outputDataCol = frame_out.dataset.map(_.getTensor[Short](iField).toDense.values(iSlot)).toArray + output.addShortColumn(nCols, outputDataCol, null) + } + case BasicType.Int => { + val outputDataCol = frame_out.dataset.map(_.getTensor[Int](iField).toDense.values(iSlot)).toArray + output.addIntColumn(nCols, outputDataCol, null) + } + case BasicType.Long => { + val outputDataCol = frame_out.dataset.map(_.getTensor[Long](iField).toDense.values(iSlot)).toArray + output.addLongColumn(nCols, outputDataCol, null) + } + case BasicType.Float => { + val outputDataCol = frame_out.dataset.map(_.getTensor[Float](iField).toDense.values(iSlot)).toArray + output.addFloatColumn(nCols, outputDataCol, null) + } + case BasicType.Double => { + val outputDataCol = frame_out.dataset.map(_.getTensor[Double](iField).toDense.values(iSlot)).toArray + output.addDoubleColumn(nCols, outputDataCol, null) + } + case BasicType.String => { + val outputDataCol = frame_out.dataset.map(_.getTensor[String](iField).toDense.values(iSlot)).toArray + output.addStringColumn(nCols, outputDataCol) + } + case BasicType.ByteString => { + val outputDataCol = frame_out.dataset.map(_.getTensor[ByteString](iField).toDense.values(iSlot).bytes).toArray + output.addBinaryColumn(nCols, outputDataCol) + } + case _ => throw new IllegalArgumentException(s"No BasicType $base") + } + + nCols += 1 + } + } + } else { + println(s"\t$name: ScalarType.$base") + + output.addColumnMetadata(nCols, name, getSqlType(base), 0, 0) + base match { + case BasicType.Boolean => { + val outputDataCol = frame_out.dataset.map(_.getBool(iField)).toArray + output.addBooleanColumn(nCols, outputDataCol, null) + } + case BasicType.Byte => { + val outputDataCol = frame_out.dataset.map(_.getByte(iField).toShort).toArray + output.addShortColumn(nCols, outputDataCol, null) + } + case BasicType.Short => { + val outputDataCol = frame_out.dataset.map(_.getShort(iField)).toArray + output.addShortColumn(nCols, outputDataCol, null) + } + case BasicType.Int => { + val outputDataCol = frame_out.dataset.map(_.getInt(iField)).toArray + output.addIntColumn(nCols, outputDataCol, null) + } + case BasicType.Long => { + val outputDataCol = frame_out.dataset.map(_.getLong(iField)).toArray + output.addLongColumn(nCols, outputDataCol, null) + } + case BasicType.Float => { + val outputDataCol = frame_out.dataset.map(_.getFloat(iField)).toArray + output.addFloatColumn(nCols, outputDataCol, null) + } + case BasicType.Double => { + val outputDataCol = frame_out.dataset.map(_.getDouble(iField)).toArray + output.addDoubleColumn(nCols, outputDataCol, null) + } + case BasicType.String => { + val outputDataCol = frame_out.dataset.map(_.getString(iField)).toArray + output.addStringColumn(nCols, outputDataCol) + } + case BasicType.ByteString => { + val outputDataCol = frame_out.dataset.map(_.getByteString(iField).bytes).toArray + output.addBinaryColumn(nCols, outputDataCol) + } + case _ => throw new IllegalArgumentException(s"No BasicType $base") + } + nCols += 1 + } + } + } + + def json2leapFrame(frame_path: String) { + println (s"run($frame_path)") + + val f = new File (frame_path) + if (f.exists () && ! f.isDirectory () ) { + // get input from file + frame_in = FrameReader (BuiltinFormats.json).read (f).get + } else { + // get input from string + frame_in = FrameReader (BuiltinFormats.json).fromBytes (frame_path.getBytes () ).get + } + } + + def leapFrame2json(frame: DefaultLeapFrame): String = { + var json_str: String = null + for(bytes <- frame.writer("ml.combust.mleap.json").toBytes(); + frame2 <- FrameReader("ml.combust.mleap.json").fromBytes(bytes)) { + json_str = new String(bytes) + assert(frame == frame2) + } + + println() + println(json_str) + + return json_str + } +} diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/java/com/microsoft/sqlserver/mleap/ScorerTest.java b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/java/com/microsoft/sqlserver/mleap/ScorerTest.java new file mode 100644 index 00000000..9f4e4455 --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/java/com/microsoft/sqlserver/mleap/ScorerTest.java @@ -0,0 +1,98 @@ +package com.microsoft.sqlserver.mleap; + +import org.junit.*; + +import java.sql.Types; +import java.util.*; + +import static org.junit.Assert.*; + +public class ScorerTest { + + private static PrimitiveDataset input = new PrimitiveDataset(); + private static LinkedHashMap params = new LinkedHashMap<>(); + private static PrimitiveDataset output; + + @BeforeClass + public static void score() { + // get model and testing data + String modelPath = "src/main/resources/adult_census_pipeline.zip"; + + int columnId = 0; + input.addColumnMetadata(columnId, "age", java.sql.Types.INTEGER, 0, 0); + input.addIntColumn(columnId, new int[]{39, 50, 38}, null); + + columnId++; + input.addColumnMetadata(columnId, "hours_per_week", java.sql.Types.INTEGER, 0, 0); + input.addIntColumn(columnId, new int[]{40, 13, 40}, null); + + columnId++; + input.addColumnMetadata(columnId, "education", Types.NVARCHAR, 0, 0); + input.addStringColumn(columnId, new String[]{"Bachelors", "Bachelors", "HS-grad"}); + + columnId++; + input.addColumnMetadata(columnId, "sex", Types.NVARCHAR, 0, 0); + input.addStringColumn(columnId, new String[]{"Male", "Male", "Male"}); + + columnId++; + input.addColumnMetadata(columnId, "income", Types.NVARCHAR, 0, 0); + input.addStringColumn(columnId, new String[]{"<=50K", "<=50K", "<=50K"}); + + // specify some params + params.put("logLevel", "INFO"); //default WARN + params.put("modelPath", modelPath); + params.put("outputFields", "prediction,probability,education,sex,income,predictedIncome"); + //params.put("outputFields", "features,education-encoded"); //SparseTensor + + // perform scoring + Scorer scorer = new Scorer(); + scorer.init("session0", 0, 1); + output = scorer.execute(input, params); + + // cleanup + scorer.cleanup(); + } + + @Test + public void outputColumnCountShouldMatch() { + // display output + int nOutputCols = output.getColumnCount(); + int nOutputFields = params.getOrDefault("outputFields", "").toString().split(",").length; + assertEquals(nOutputFields + 1, nOutputCols); // "probability" is a vector field of size 2 in this case + } + + @Test + public void outputColumnNamesShouldMatch() { + int nOutputCols = output.getColumnCount(); + + Set columnNames = new HashSet<>(); + for (int iCol = 0; iCol < nOutputCols; iCol++) { + columnNames.add(output.getColumnName(iCol)); + } + Set fieldNames = new HashSet<>(Arrays.asList("prediction","probability0","probability1","education","sex","income","predictedIncome")); + assertEquals(fieldNames, columnNames); + } + + @Test + public void outputColumnValuesShouldMatch() { + String columnName = "education"; + + int outputIndex = output.getColumnIndex(columnName); + String[] outputStringColumn = output.getStringColumn(outputIndex); + + int inputIndex = input.getColumnIndex(columnName); + String[] inputStringColumn = input.getStringColumn(inputIndex); + + assertArrayEquals(inputStringColumn, outputStringColumn); + } + + @Test + public void outputRowCountShouldMatch() { + int rowCount0 = input.getRowCount(0); + int[] rowCounts = output.getRowCounts(); + + for (int rowCount: rowCounts) { + assertEquals(rowCount, rowCount0); + } + } +} diff --git a/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/scala/com/microsoft/sqlserver/mleap/PredictorTest.scala b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/scala/com/microsoft/sqlserver/mleap/PredictorTest.scala new file mode 100644 index 00000000..b98ddcbe --- /dev/null +++ b/samples/features/sql-big-data-cluster/spark/mleap_sql/mssql-mleap-app/src/test/scala/com/microsoft/sqlserver/mleap/PredictorTest.scala @@ -0,0 +1,112 @@ +package com.microsoft.sqlserver.mleap + +import java.sql.Types + +import org.scalatest.fixture + +class PredictorTest extends fixture.FlatSpec { + + case class FixtureParam(input: PrimitiveDataset, scorer: Predictor, output: PrimitiveDataset) + + def withFixture(test: OneArgTest) = { + val input: PrimitiveDataset = new PrimitiveDataset + val scorer = new Predictor + val output: PrimitiveDataset = new PrimitiveDataset + val theFixture = FixtureParam(input, scorer, output) + + try { + var columnId = 0 + input.addColumnMetadata(columnId, "age", java.sql.Types.INTEGER, 0, 0) + input.addIntColumn(columnId, Array[Int](39, 50, 38), null) + + columnId += 1 + input.addColumnMetadata(columnId, "hours_per_week", java.sql.Types.INTEGER, 0, 0) + input.addIntColumn(columnId, Array[Int](40, 13, 40), null) + + columnId += 1 + input.addColumnMetadata(columnId, "education", Types.NVARCHAR, 0, 0) + input.addStringColumn(columnId, Array[String]("Bachelors", "Bachelors", "HS-grad")) + + columnId += 1 + input.addColumnMetadata(columnId, "sex", Types.NVARCHAR, 0, 0) + input.addStringColumn(columnId, Array[String]("Male", "Male", "Male")) + + columnId += 1 + input.addColumnMetadata(columnId, "income", Types.NVARCHAR, 0, 0) + input.addStringColumn(columnId, Array[String]("<=50K", "<=50K", "<=50K")) + + scorer.primitiveDataset2leapFrame(input) + scorer.frame_out = scorer.frame_in + scorer.leapFrame2primitiveDataset(output) + + withFixture(test.toNoArgTest(theFixture)) // "loan" the fixture to the test + } + finally () // clean up the fixture, nothing in this case + } + + "A Predictor" should "be able to convert PrimitiveDataset to DefaultLeapFrame with same field names" in { f => + + val columnNames = f.input.getColumnNames() + val fieldNames = f.scorer.frame_in.schema.fields.map(_.name).toArray + assert(fieldNames.deep == columnNames.deep) + } + + it should "be able to convert DefaultLeapFrame to PrimitiveDataset with same column names" in { f => + + val columnNames = f.output.getColumnNames() + val fieldNames = f.scorer.frame_out.schema.fields.map(_.name).toArray + assert(fieldNames.deep == columnNames.deep) + } + + it should "be able to convert PrimitiveDataset to DefaultLeapFrame with same int field values" in { f => + + val columnName = "age" + val columnIndex = f.input.getColumnIndex(columnName) + val columnValues = f.input.getIntColumn(columnIndex) + + val fieldNames = f.scorer.frame_in.schema.fields.map(_.name).toArray + val iField = fieldNames.indexOf(columnName) + val fieldValues = f.scorer.frame_in.dataset.map(_.getInt(iField)).toArray + + assert(fieldValues.deep == columnValues.deep) + } + + it should "be able to convert DefaultLeapFrame to PrimitiveDataset with same int column values" in { f => + + val columnName = "age" + val columnIndex = f.input.getColumnIndex(columnName) + val columnValues = f.input.getIntColumn(columnIndex) + + val fieldNames = f.scorer.frame_out.schema.fields.map(_.name).toArray + val iField = fieldNames.indexOf(columnName) + val fieldValues = f.scorer.frame_out.dataset.map(_.getInt(iField)).toArray + + assert(fieldValues.deep == columnValues.deep) + } + + it should "be able to convert PrimitiveDataset to DefaultLeapFrame with same string field values" in { f => + + val columnName = "education" + val columnIndex = f.input.getColumnIndex(columnName) + val columnValues = f.input.getStringColumn(columnIndex) + + val fieldNames = f.scorer.frame_in.schema.fields.map(_.name).toArray + val iField = fieldNames.indexOf(columnName) + val fieldValues = f.scorer.frame_in.dataset.map(_.getString(iField)).toArray + + assert(fieldValues.deep == columnValues.deep) + } + + it should "be able to convert DefaultLeapFrame to PrimitiveDataset with same string column values" in { f => + + val columnName = "education" + val columnIndex = f.input.getColumnIndex(columnName) + val columnValues = f.input.getStringColumn(columnIndex) + + val fieldNames = f.scorer.frame_out.schema.fields.map(_.name).toArray + val iField = fieldNames.indexOf(columnName) + val fieldValues = f.scorer.frame_out.dataset.map(_.getString(iField)).toArray + + assert(fieldValues.deep == columnValues.deep) + } +}