From db5f75ae14e6d6d79c3ca2904807ecdf340fdde9 Mon Sep 17 00:00:00 2001 From: Vittoria Tommasini Date: Thu, 20 Aug 2026 07:28:06 -0700 Subject: [PATCH] Add option to utilize GPR model to create zero eccentricity initial orbital parameters --- pyproject.toml | 2 + .../Examples/InitialOrbitalParameters.ipynb | 74 ++++ .../Examples/gpr_model_adot.pth | Bin 0 -> 13374 bytes .../Examples/gpr_model_omega.pth | Bin 0 -> 13397 bytes .../InitialOrbitalParameters.py | 404 +++++++++++++++++- .../Test_InitialOrbitalParameters.py | 207 ++++++++- tests/conftest.py | 21 + 7 files changed, 689 insertions(+), 19 deletions(-) create mode 100644 src/SimulationSupport/EccentricityControl/Examples/InitialOrbitalParameters.ipynb create mode 100644 src/SimulationSupport/EccentricityControl/Examples/gpr_model_adot.pth create mode 100644 src/SimulationSupport/EccentricityControl/Examples/gpr_model_omega.pth mode change 100644 => 100755 src/SimulationSupport/EccentricityControl/InitialOrbitalParameters.py create mode 100644 tests/conftest.py diff --git a/pyproject.toml b/pyproject.toml index a78c042..ad28e17 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -19,6 +19,8 @@ dependencies = [ "pandas", "sxs", "varpro @ git+https://github.com/sxs-collaboration/varpro.git@978106eaf3d7a6a7f0c5f167726d8e0fc59fc95d", + "click", + "rich", ] [project.optional-dependencies] diff --git a/src/SimulationSupport/EccentricityControl/Examples/InitialOrbitalParameters.ipynb b/src/SimulationSupport/EccentricityControl/Examples/InitialOrbitalParameters.ipynb new file mode 100644 index 0000000..19a13a2 --- /dev/null +++ b/src/SimulationSupport/EccentricityControl/Examples/InitialOrbitalParameters.ipynb @@ -0,0 +1,74 @@ +{ + "cells": [ + { + "cell_type": "code", + "execution_count": null, + "id": "0c6702c4", + "metadata": {}, + "outputs": [], + "source": [ + "from pathlib import Path\n", + "\n", + "from SimulationSupport.EccentricityControl.InitialOrbitalParameters import (\n", + " initial_orbital_parameters,\n", + ")\n", + "\n", + "CHECKPOINT_DIR = Path.cwd() / \"Examples\"\n", + "\n", + "target_params = {\n", + " \"MassRatio\": 8.0,\n", + " \"DimensionlessSpinA\": [3.61e-14, -4.5e-15, 0.-0.8],\n", + " \"DimensionlessSpinB\": [-1.1218e-12, -3.82e-14, -0.8],\n", + " \"Eccentricity\": 0.0,\n", + "}\n", + "\n", + "D0, Omega0, Adot0 = initial_orbital_parameters(\n", + " target_params,\n", + " separation = 13.378235,\n", + " method= \"PN\",\n", + ")\n", + "print(\"PN:\", D0, Omega0, Adot0)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "443ad8b9", + "metadata": {}, + "outputs": [], + "source": [ + "D0, Omega0, Adot0 = initial_orbital_parameters(\n", + " target_params,\n", + " separation = 13.378235,\n", + " method = \"GPR\",\n", + " gpr_checkpoints={\n", + " \"Omega0\": str(CHECKPOINT_DIR / \"gpr_model_omega.pth\"),\n", + " \"Adot0\": str(CHECKPOINT_DIR / \"gpr_model_adot.pth\"),\n", + " },\n", + ")\n", + "print(\"GPR:\", D0, Omega0, Adot0)" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "env311 (3.11.6)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.11.6" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/src/SimulationSupport/EccentricityControl/Examples/gpr_model_adot.pth b/src/SimulationSupport/EccentricityControl/Examples/gpr_model_adot.pth new file mode 100644 index 0000000000000000000000000000000000000000..24115aaf5700b62024b9119243c86d3c87a9f2b9 GIT binary patch literal 13374 zcmcgy33wAl``bK`ja@ zQbgsFo6kSu0SdN3nnY9-IRve8sEC3GD4!GsJUG6Y*(BYx-C8Qof1hWM%*^}yz4N~B z&b~8~9ChnpiV6v#8XOT+2dc0bYKtrddXrXHU=c56`ff>uE%=(PpLd;zq|&$pON`g{SkvW8~E0zK3hj7Bw~pmj74 zf}=E%NX%rMt2Y@97E3{l*Y#Ga5nnnjHelU43!n=Rx@wY}p&(M&Y!#r+Xcl5jmU()p&9m6d1)!S&sby$q z5NYE;cLSm8f#{;6MVc@Hg!7>1ZPVC_i~Tea0W=EGOQh*7q3I)lzC7rs>4SO$^JZ)3 z>Y-V0A|^KJ&4q%&ny)kILH_`xcrXAV1q?JRg9FGCz`df{LnLZP2>`j`zS|&>+GS_}K_VC?5)79R#0X#n4`Ma_Q4fkx*g~_M zY2pIFl7K=qO}yVU38HBdMbjh+AejeBO~O9`%UOpG;7bN6qVgjp%BKn-jR&gRl#a7c_4Xi@x-z_HkbVw)fa&X@os8f1t_nG&S20?6XQIL*Dt zA`+xL9eUwq!cpx*HFyJYyx-arL~GwKT02_+IXsxCNozK?oY{Flr8gWWNxCpu08@A{ zwS_cN%Rk_!@dl+vLNiSOxjcA~M7_jPcs!KYT&@SxebAa#$Qp66V20?6nGy!A0A}$( zcN+{+@t{Z(UFTmf zT2tX;_8}X*FRKMXs)O)2RxE}O5W@!wz{&$b!`=?Ooc(QnR_`lYBI(FH0X)Qm(v~wx z{V?CpBawEYJdh*- zED}kckdQnnfWxHkG}yd%~o3C(5!Y~jJy7Sc$2^B+Hrci4Ja zLi2wD*v5lb!|H;NcU0Rap?OaL@AKe;7Sc$2v)@nS9oIgT(0n9-13dUx(*unUMS7jtH>}Om z8w(AB6&&;livw^XIE2PVl#L%o7^3r!pin^B_$Oiv92LMZ9(*d!#3XomMx7OW76_IE zYQ%x|a}0)NJzwCzdN&k{#5FV)3^jCROUP9sDLDVDHxDq2PBrxFS|y@77gFM>25L z4z7uf@Z_-ZMosDehaLPW%BQwOUfPkr?BH(^3OiJyiu`v8T(^T8q6~Iuvt*={H|^k- zD2^TK6&L*qDYOJ(he7Bnw&-2Ih18C)6=~ht-x|*+1MsA+n=mL0CO~ccKuimPB;AB< z31&M#CXUcR<>XBaLx^yDzc7x|<_b%d2qicj{5TyEjx-U09Z7>u{sv+78u;=^*qJo! z;%|t>9`Y9kkf&+-yU4AiH zKLv<2<)=u3c()%>)=$kvlp5q70zS|WF6*bJ@*_BipbYk-$oi=%6scPG62KvTzz)dq zl0*kak@{#~z0I1Z2}e!p5(#bHQZl>`H$w}va433+A`N|sLBL`$9EKZu*IfakVa(^k`(k)fFseBhpFglv%@rWGmx|ms?cMU z;j|qxxQyg&*dvmjLKfB7A?L4;46G*s?rZN$$6z`U8s!(l>71`ofNP}YtINdTXd*wx zCy&D*1VQ+cPMAUJGto6l(-~D5&1l(CJ6o?4Y*257V^PHnm_@3)lWL=2)S0waeX$Pe z@XiAqhblGFx+0y`swI_BjqC98uhA^Z6&Q=qwgM6}=}{wVvC$l-oe#(Vb29~MHUTx$ zYp_|xl{&3iSEPscV_~%ELg76Xn2mqoCS>M^Hb`tmW|)HtXTXWL3a1RkW^JJj^Aiak z73(zJP`RLkg?d3-pr1|DMSYuu>g;ebx{4bhIJvc%bZ`oW!mkswnIlVdCWIC_JTmA0 zNP-Ke;%dBygQ}5VOz;6*fWMtSI}m2PUor>lHV5lqvk~RYh(c#F&KHe44OQ4-F1kjM zd4T91{GP)HQQ-_Y9am+t!}4F~FspKOCfjMv)_2e5{E+E>DS7(t&woCe$*eBD@!j#* zGCi4Xi1JpU0L~!dGl@9!TUb?fsrpJL8yr3F`l778natTGUoIG%d^M9f7yMJV9y5Ok zush+`BGF8yN71!~CCcB3a$c`s^jV}bR(bZ_gSK$7V`K}*XfD*^$qv+^26mW->B+W{ z2Iiw@-zz9i%93FLE(@Jt2oj(kJ>pR&7p*YD+2{wFp5-orh3H8fX$UJS-Y!c>e z=u1)f%#q#i>qYgczu%WvD(EeEn=T{0r*Fg0yYJ__nT (qoqLw6ne2qKLt#JWDv0 zT%QG-cSLMPdb|vUkE+r8zU=cJ%wFmoupGZr_-u-uZj?PPBR#Pt?76uSmB*Br-x9WG z?LYJPq?C>GqwXiPNCDJc)Zca?{k|D}_rKsr252e2qYU%^HPn9UU+^nh%3uCv1o8iO zZT>j-FZknI%72IKpU`stJ7WLDmh!h98OC%nbYZgZ%3*elsbX*}gsne8<39l4Q~q7- zNeyfKT}aAEk81Y25azj*hB=jkMi2ZK{!eac|JhDcIqDYAd5Uv5cGVP)a_}70=RwYM z7MhfqYlsHJg*gTjtf&c_HD!m=5W+o1)O8RgB-O1@h;bMRN{Qj zVVuCR$26R05R$z(gL5p=a_p8-NH@Sa)SX8eAHLSITp<1*jy9pKWfuwAxC{-;ixGH=aIP_wcX4)V&-z}>W4YX zGM{4?mT_#Eg<~fRoTD1aR6x%2ybbA(WR_bw$1kNE3MF!uUHg-*eCj59>((tc_RJBs z<0tR2xUP2SMK+;g2m9=mYveaBTl2#n^1R@{Q|wd6K4m9T2iTiQ@3KY5-eC_IKVhfC zU8Mg0w?Ae7IJBFs3URTg7d_90e3QlY=Ekw35>B(9UoK}yCOpdW#}={+yBpazD|)iq z!}_y-BmMZgg%Y<%bWAcjI&a_P*r)TUj6XV&2%9>RuB+p80UVICs}QeNzzzbLc%v_4GU?^Lm#t zOeXBbR5=sfNmq`lmPRbBKH<3O!aSYw)0hb>ueuJQ=UPRDt8mT+6_zFPsEm5;po`k6 zVyvYdm}Rbo)p2>Fm|b0Oq=mfoJ(1siz8@2IaZ2^xwOQ^be^4;3E_89Hzu%h3Unoy; z7eAUo^c}ZS!JL{zGnfbS`Z!)(-kJPj>~bXgascX1bY?2=t7&j+WlF78qL zUUro&TBWK8{!DfCj9O*+^OEbMHg~65PF!~Vwn0!nNV>!>Xs=-&Y;^_XSm(H-nh@Uq7!pb-0xD7wLlaVVJf*yYBszbst1rS>uK= zwW}DWEPJH;MpV3e;Hpgb#b;Re%&#-tuTyHKq-3(HqFT*-uxD`^)@^mK2o=W2D%tEZ zpI@umcldW#;@Kou_Ni3&{H%1+&m%u(5S|cCTV{!VhwBKE!*-ZzJmrcyG1k5M!btb= zb9+;MGb>16P(8UXUVfTz405$%x-BeGor&A(+EX%`*bsB#es4(|%7ouDj^JH-`C@g+ zn^_FT$GYKqd}gYHh#fbrSm(N#peAJ#cV)OXrLI(Mz5B3h*68-t-v+g2T+WdMs z^F`@Q^i+j=dz2dEGS8-d=$gDbgc&#OLh2(I?qhau8|xmw^bJ+5Fy6H>bXS^jih?=Y zePbH@xf7$^vf0(9_&HTs)r-}Srys4x--eSf#H`!U2Iar#Z0dW|ukr7ZG9$g0?Dxo! zywllT6`XQ~?U-_7>XEAco1S!TDc$INrspymd>ociOV4)}CK{Z_J1lj6d#};irTA&* z#SXnvPOjXh+*R2vW$oY{$%;PDIqw^*pxIR+DI+h(D!(nurJvuHPzlBlt?bt!&Dno9 za2ooaq|X^2r5AkBD&^!gKyQ4|U9}>Av2#Jb#kuo|xk}rt@~UN74>{LedPuqV_=of% z%P;i6^)>W=Cf3r6ul=qB=SrO+TC;M0;b8h~g|afZqBETx^lDW>!V#FZQT3R*a?=H@DGjiyoc|R$S9hst>b)r^z`B$|v=cRUZc}cW# zS9*CR_T#gv%r$t@gim=thEw5<`Vpm=-5R#~2&pTeo(jrJ3zlF<$O+34g66{7DXvSoAU4|4rS>0=A~e`oBE zQpOhY=cyiI2RM7PLtOh9P4s*g+hXPKaCgh?k5bA1cl)Cus)hZ(Gxn!j%pdK!-;Hhl zi2e#$6T>;Y{JY)A{Mq}jx2{KBT#CU~HlyIUy;IPEMxd5*10bGER{khb{?9C9OzKOo)w-rIQtOEY@;!*G@z5 zrW!tT)InHT|39~-YdjISBH)nuA9sBa=@L$_0Dwzq>8Kuve{rZjigYOkFFDp7Q~NbQ z5f^{zLr7Os0wHQ-AmZvnNEbx#mQdX>8LjWu>kH8$4jDPnsM1A^K#UqWj0Sp07bXHR zCJ&I&Ls4H3`M%W1VZ_(RknhXnyJYl8sE;9?*a!BdMh+vfK8AFn9*8k3Qbv!Y`WVuL zE)b(e4kNifhBTQ7#JGh1liX;0poKp9_y`QsD)f#vhS9*a(rT&eYhw@%JS&Zk zK!^=85DgqF4URyFu=^U_`~e3T#XZx_(?pKPQ2)wW1+}_XAcRLps>wf2!0S-G?_&qw&vBj=F7IynBg{0j-gg s_waomeD^_k$sQe5-X=&6?*&m`;Xh%Z1Et&Od((!ril8DW++J+^KiZ-m8vp5YbBv$-(ZWHwmz+5&T#4n{-Wd@cD6Iu`I|SWq0Twekj|B}5DL z`BsCmP|NF07BkeADMCP}NKqvV>ZvQzgU+HzNQTj@<0p$ZLqM=50P)3uE}B5}T~KNT zA?g4g+&4*0sRNdQ`|+u3MKRW`8xOj3AXKAlg@Q<7lZA&mgNcten&<1GHs5SD6@nhc zNUh?|#UgDi=vhqY9zb-%hKV#`Ja~`;z1mD;EiLuYg!|DbKyQ(zkA$W#5BhPSzosv; z2IkMv&eKDa-bh?*)SHUo~iUTp40my<9 zl(x_$=bBhQutcB`T@&YXO}yxu1kp8#JWz5VNfZA^z;fP6_Ty86$3*2*B+94qAdLg6 zwsA-uMEh|h0Y>CtB^>EI7|nq(BymWN6lziaki;=t`BEz{CeBztq-2mGB4tXD#_=GF z1KFCP$RiS@d>wk=WztdYMJ@IuVvf(-<3(>z5WPK-2e}-Wq)BTvww&8JAEhT9CreD2 z!h<{xOl>EP)boGx(Rh+lBcYkbgXtW2oMgSkQ+Pa-_*~9{8D3~j2jq>|7%)?e#ViSf zmIt#rplbtzRQY^=hFDM_G89S}^gNiufgPfC~;@xaD`#bj`3F#vLQlzY*|Jcx6r5Co?GZ6hng13qX&vM{7 zjRHkm2Ct>8`l}DyGbmLeY%!}=i&^zIF{{?_U@Zrp*JQPeM=mVueDt0H>II2uFY;hL z2R7UTN@+-o4+a%X^`0zXf*FPnAZ}8wv4(ycAwbj@?aka_R9ypMIO0atMbu%hO%mjX#x)%9C-U4P)egyM9X|)ncCl6d4 zaLXq_ljz7<_KuInGpZes&>ZB!yBs*wP8z8-hkZ1jaqT?`&HFq!!hsJo51{d(M6WY> zhqd{7Ls2nr0Y|;U;sBfgj-jy;RpZAIh8X-0Q7WKn{39_3KIXv*4xAKcViLT3gU$jz z@drx;b>cw#DF#Ecp3m^qrf{OMx>s zaJCp}V9|u2pAEA(aLxw45gU5>HAts-;9oXyUhEVl+lf4E-Z`OniUq#4f$v1g#I{LB zYtbalWSL`zCE$V$eBYpa!<-2G+Xj9R8y-RE4GlyW$z(~ATHvA$ToPqsS}P;f&6UD)R3o@VgD%7G-ciTO}i< zyki6PqBstyM_i06q|g$C4F;gI*rQM59#TKT4y1RWuQ#4g`r%1?x8P71M1VT_farDt zNz8XxqgNbk#pD@nT)(T6N2q8H4`Ec$>IMPA{-cLGo^>qkq)WKUv z!fvEvcV9=G-clrQ$LJUsO1k#&bsZ$zRT>}IlZZXw6T|b><{6}7Erm6VAU^0r9Cx>f zQrGq(;Nd=S{3y$)+d`mw6O=wa6#PK98AYmAUjo?A2Phk-SlFL*9pLLK8z;Z|2o5A- z5BbDoyrhC)Kyu#^mk<8EkS7Dl0qNYc?;82Bv~!x6ZnXWg~9 z3#oG8NOVsubYL{PM70hag}zC-17pw?4`b1pg9>!kj6g~>i!FKOOV@c}9Fp=;ii^B3 z9^Gpw_Fw|~$-_i+=AaUtTSp?PC1Dbh@{x#3yf7KvlaM@ye)2E{ojI6_&Q=>tLtn+D zY=bIv8CiVB25Eec)NME-5=$Y2T5OQ@wMY23lLYtn_m*QYod}Kg3E^_i+sQ99Quj4f zVsH$RAM2IJX%LJcyhSI>Anlpx9I5Gs8Vn}1?5Lfi*YQ@Ux4?0zVJ6HX%{@u8fj8)k zT8q9^2X%Po0cN8{jkK*qXR&BWBUIxyy!>l0iE@R8611&=#Eg2>$x>=C#cCJAoImcS zK;6cpZh8%NtGH69HR($9Z~_+Yh;QOO6*v+9#C^!r5ABdxOH42q-JA(0;U;WksYzR8 z#R`Zb9J;U5^gs=~4i@QoZJ~Y+Q5ab_8MWEq6m%B1KyZ0$HR@m<#=;L2w3;HybVdXh zF)|`|LIeSZQ*kri#X-#|GDi3(d;?b_UQd`|zJW)coCEd1sy7uZfOz|4F4lA|R>*4A z;PS?V_;g0YLebaLP=gIlN9RZ~EfAxHAA9&Xx;Ychz)hLVh=S+un_V>~lj*wQm3PkP zUdVL4q@1zu(@P&`((4!8{@1CPay^-Ri1HS`5Y8mxvxqqTb69ol0`>JwCTQ6F-(P!gEdSuWCU+q2T z6b&OgItEjb7SDa44t21>d`wSvkud_@0bf%*PM>q%l zK(n>!%V80^5(l1Qbj3phI&*L?IuqDwa2~q!hfRV;d@n>jorfjpil8ooCUl91GBf%X z5lhh(4*@!J5Tdiy8l9`|9DFQ1Ztl*06orpm+3mkR&20bWFW~iNylvO8|7SGtyit!~ zm!+NW(+MR2KICD0QHjm&DvpZDNLP6F6vdsz)}3d#Pq{IBbn(bSw2?{Sqb2sF=GkL2 z(&O8~K7D$4)rlm`Zw_;8{A2!bN-4vS%#Ul20?1rsZ|C6tXU9BzH~h!}?c}$YWB&g} z*-qaLzoMP|{pWq{#QcL3Gm7qpKdznp_sakA?dHEX{!eHpzjN~kHJ=>2S#$XE<(iFw zMI;s#q1;CN2LOD?Ka1g1l+3e`n34WyYoCR{cdxr}{8wI2Uw1e0SGLoBrt4Ies^?gD zX)epWn8#9fj-~oO&brT|V_6DIP0VANefg~0n!~#Lq@yE)Wj2jw?VTsH_98Xwc1~xh z1xS`+S@)ykS@$W1b)Ol{QiUljb7LCIe3Qd6j>#-_d=~3|qlC3zoXRr$5g(Y#Qg0Qq z_JNPH)E30MTF+7m3t78iJjavlwUn*tY@8q-8jRMyFyPchVZ3*i>HlMYJ7PIz0^X$2om$LR9Wh~Rpz}naAS^MYZ ztUIQhweL2t?g}ew&!35SO)S+5VZ1SorFt$zwnEmvXco&%HzNB{pT0|2>N7J-5pVqTUhs{s&$jqh!^;2_tFj;v%>o;b-&qK+_k z65nA;PQ1+=F?__#fO|>%g#9O(Uyto$s)L=(ndNJl;ImmwA2ypA9e;-T^rvM^O8ioW zJF$dW($m1~sqDq<2phoMLi~$%)691VCNdj&3-if_Q05N0cjlQSW=}yRQA8>>6@!)xf97_Z`H6@8cu=+1qrNU9b5#w+b(l`^S&vxbSg!M?doY zb9rCalwpzN`s(hRs>#a!V*4ZGo!_pmRMqJ>r@b*FlHR#H)wLlmhkkEhwky=pFBM_X z#|~1ezsygf-|Rk?&V;?`Y9Yauc>QD5v*Al>Keyj;VxF!AY4rFPZa9yj>uZWiXVKg( zDlALn(LL&oqfTnKinc7ckN&H3No{QYXnJqx?X=*xz9aJce&|n!UCyhm*_h>e=7NIm zaIw29{oOz!e{tDkuF|C$MBnTe6!huI$u#D{yk3rVtGkikFO5o}Cw&p-8dlz&?o-;` z6|-Wm%6;fG*2OjYz^l&k-!Y?;fa)xOGw0aZdof@awgz)5jMOdyy_!ADZrT z$vN~S<2s7Cva(0e*IuOQ@`)*~+mUgu!7paIF0W!-v%b!7y-BI*va%_v%33vj$Wf7o zbz9#%T!k^R%3gMwe%PowaQqi%!udq!#M7y+g<0vu&JVxOAUwgEPV@@>F6W0x4*OxM z;j}aI^Kq{A7gJoPzO8xeXOn{1g4)S<#p*MJV~Dc@-D63a>RjwAPDj}o;zP`d?cSai zML#$wo8Vo2^>S_5o-7*UW8H8&J~Gt-#E;w7Y;xX-SCe~__GUP@rM{qg<>BMb*<-rY zeiIN#J7>^zmwqnwJ?E75!F2Yti>Zq*4yX6MHqMpv>`v7+ zKF7H=WN(@wPeGsWxit-5>Pl<3zwGQ-`Zra1^}5=p(m$@n?}jg)kKVKeO}OyU()*~N z%=;)QBRyR1eRN}7H%Da0L|2o(6nsm_%kJ;6*FNBQx_+=jd5}%|k-H*1RI^RE8aKK6$tzDQtA>O+PR0Z~qQ-S{ zR9^^j1a|mcD0=9Qz^&`z_#t>5Q^Mi3-51O zCMe@V9ZcReVT&ugdi$4k)%mrns_kf^>Jeoph5GF0g{{Wp!h*@`g+qPT3MmsSgiYECp~LnHVQ6qkRs7Cg!ks71R8vm` z3x;RURvlQnstU)k;Hy=2c$S0@`8f8X!kUjGDltQwgb|H}ZKz*+Q8M!fR6e$E{&42d z8ix9QC3AZG>x?xbT!rlnuX%U%AJl$mm)0(i#dg z(ucP435-clDB}}jlak z@R6$yz{>jmxh-AP>5UuwK63x#Zj2&b!@+BR4WD|nSkxHBw;a?MM!Fb-w__SUQ~NiC zAs&orj3Hf6@yDo>!$3~*cZzgTB)mZhbTtL7@isPx=n>mQ5$Vc?KS-S{NE1z@%M<<} zQwGXuqG-K2b+RCFjX~tinevdFCh?6yq%(YfbLwP45*mX@r|kY9vm@j*No))vt?m3l z>SRHbjX|V^g+Itu^#A1M^8~H>`6rvyvHl5LjUL!$L7GHZ`sVk?2t$8%Z5E?RfTc8xSriwKGa>`YA>V!8n@n>Kl67zh# qN_oFi@Jc@VRC~H0IlMtc{fz&cfj+1TN3Wrdq*pi Tuple[float, float, float]: - r"""Estimate initial orbital parameters from a Post-Newtonian approximation. + r"""Estimate initial orbital parameters from PN or GPR. + + Estimates initial orbital parameters from a Post-Newtonian (PN) approximation, + or from a Gaussian Process Regression (GPR) correction to the PN approximation. Given the target eccentricity and one other orbital parameter, this - routine estimates the initial separation ``D``, orbital angular velocity + routine estimates the initial separation ``D_0``, orbital angular velocity ``Omega_0``, and radial expansion velocity ``adot_0`` for a binary system. The resulting parameters can be fed into an eccentricity control loop to refine the starting parameters. @@ -46,11 +56,25 @@ def initial_orbital_parameters( * ``"NumOrbits"``: Desired number of inspiral orbits until merger. * ``"TimeToMerger"``: Desired time to merger. separation : float, optional - Coordinate separation ``D`` of the black holes. + Coordinate separation ``D_0`` of the black holes. orbital_angular_velocity : float, optional Orbital angular velocity ``Omega_0``. radial_expansion_velocity : float, optional Radial expansion velocity ``adot_0``. + method: str, optional + Either ``PN`` to compute parameters from the PN approximation, + or ``GPR`` to apply a learned correction from a trained GPR model + to the PN approximation + gpr_checkpoints: dict, optional + Required when ``method = GPR``. Maps the quantities to correct to the + path of their trained GPR checkpoint file produced by ``save_gpr_checkpoint``. + Recognized keys are ``Omega0`` and ``Adot0``. Any quantity without + an entry is left at its PN value and no correction is applied. ``D_0`` is + currently never corrected and is always returned at its PN value. Each checkpoint + is checked against the quantity it is supplied for, so passing a checkpoint trained + for a different quantity raises a ``ValueError`` instead of silently producing + a wrong correction. + Returns ------- @@ -66,6 +90,9 @@ def initial_orbital_parameters( ): return separation, orbital_angular_velocity, radial_expansion_velocity + # Unpack the pieces of target_params we need. Everything is derived from this dict + # instead of passed as individual arguments, so callers only have to build one + # dict per system mass_ratio = target_params["MassRatio"] dimensionless_spin_a = np.asarray(target_params["DimensionlessSpinA"]) dimensionless_spin_b = np.asarray(target_params["DimensionlessSpinB"]) @@ -99,19 +126,76 @@ def initial_orbital_parameters( " parameters: 'separation', 'orbital_angular_velocity', 'num_orbits'," " 'time_to_merger'." ) + assert method in ( + "PN", + "GPR", + ), f"Unknown method '{method}'. Choose either 'PN' or 'GPR'." + + # GPR method. This will be modified later to accept + # both non-eccentric and eccentric GPR models. + if method == "GPR": + assert gpr_checkpoints, ( + "The GPR method requires a 'gpr_checkpoints' dict mapping the" + " quantities to correct to their trained checkpoint file paths." + ) + return _initial_orbital_parameters_gpr( + mass_ratio=mass_ratio, + dimensionless_spin_a=dimensionless_spin_a, + dimensionless_spin_b=dimensionless_spin_b, + eccentricity=eccentricity, + separation=separation, + orbital_angular_velocity=orbital_angular_velocity, + num_orbits=num_orbits, + time_to_merger=time_to_merger, + gpr_checkpoints=gpr_checkpoints, + ) + + # Compute the initial orbital parameters from the Post-Newtonian approximation. + return _initial_orbital_parameters_pn( + mass_ratio=mass_ratio, + dimensionless_spin_a=dimensionless_spin_a, + dimensionless_spin_b=dimensionless_spin_b, + eccentricity=eccentricity, + separation=separation, + orbital_angular_velocity=orbital_angular_velocity, + num_orbits=num_orbits, + time_to_merger=time_to_merger, + ) + + +def _initial_orbital_parameters_pn( + mass_ratio, + dimensionless_spin_a, + dimensionless_spin_b, + eccentricity, + separation, + orbital_angular_velocity, + num_orbits, + time_to_merger, +) -> Tuple[float, float, float]: + """Zero-eccentricity initial orbital parameters from the PN approximation. + + This is the PN-only implementation, which also serves as the baseline guess + the GPR models learn to correct.""" - # Import functions from SpEC. These functions currently work only for zero - # eccentricity. We will need to generalize this for eccentric orbits. assert eccentricity == 0.0, ( - "Initial orbital parameters can currently only be computed for zero" - " eccentricity." + "Initial orbital parameters from PN can currently only be computed for" + " zero eccentricity." ) + + # Import functions from SpEC. These functions currently work only for zero + # eccentricity. We will need to generalize this for eccentric orbits. # These functions call old Fortran code (LSODA) through # scipy.integrate.odeint, which leads to lots of noise in stdout. We should # modernize them to use scipy.integrate.solve_ivp. - from .ZeroEccParamsFromPN import nOrbitsAndTotalTime, omegaAndAdot + from SimulationSupport.EccentricityControl.ZeroEccParamsFromPN import ( + nOrbitsAndTotalTime, + omegaAndAdot, + ) - # Find an omega0 that gives the right number of orbits or time to merger + # If the caller specifies a desired number of orbits or time to merger + # instead of an orbital angular velocity, root-find for the omega0 that + # produces it. if num_orbits is not None or time_to_merger is not None: opt_result = minimize( lambda x: ( @@ -139,7 +223,8 @@ def initial_orbital_parameters( f"Found orbital angular velocity: {orbital_angular_velocity}" ) - # Find the separation that gives the desired orbital angular velocity + # Given an orbital angular velocity, either passed in directly, or solved for above, + # root-find for the coordinate separation that produces it under the PN approximation if orbital_angular_velocity is not None: opt_result = minimize( lambda x: abs( @@ -163,13 +248,14 @@ def initial_orbital_parameters( separation = opt_result.x[0] logger.debug(f"Found initial separation: {separation}") - # Find the radial expansion velocity + # Now that we have a separation, either passed in directly, or solved for above, + # find the radial expansion velocity at that separation new_orbital_angular_velocity, radial_expansion_velocity = omegaAndAdot( r=separation, q=mass_ratio, chiA=dimensionless_spin_a, chiB=dimensionless_spin_b, - rPrime0=1.0, # Choice also made in SpEC + rPrime0=1.0, ) if orbital_angular_velocity is None: orbital_angular_velocity = new_orbital_angular_velocity @@ -193,3 +279,297 @@ def initial_orbital_parameters( f" {num_orbits:g}. Time to merger: {time_to_merger:g} M." ) return separation, orbital_angular_velocity, radial_expansion_velocity + + +# Map the user-facing quantity names to the 'output_name' stored inside the +# trained GPR checkpoints. The checkpoint names are decided within the training +# pipelines, so we map them here instead of renaming them. This can be changed in the future. +_CHECKPOINT_OUTPUT_NAMES = {"Omega0": "omega", "Adot0": "adot"} + + +def _apply_gpr_correction( + quantity_name, baseline_value, available_values, checkpoint_path +): + """Load a GPR checkpoint and add its predicted correction to a PN baseline. + + Note: GPR checkpoints predict a correction to the PN baseline, not the direct + quantity itself. + """ + from SimulationSupport.gpr import ( + load_gpr_checkpoint, + predict_with_gpr_model, + ) + + model, likelihood, meta = load_gpr_checkpoint(checkpoint_path) + # Guard against a checkpoint being passed for the wrong quantity, which + # would add the wrong delta and silently produce an incorrect number. + expected_output = _CHECKPOINT_OUTPUT_NAMES[quantity_name] + if meta["output_name"] != expected_output: + raise ValueError( + f"Checkpoint '{checkpoint_path}' was trained to predict" + f" '{meta['output_name']}', but it is being applied to" + f" '{quantity_name}', which expects a checkpoint predicting" + f" '{expected_output}'. Check that the checkpoint matches the" + " quantity." + ) + + # Assemble a raw feature array, in the order the GPR checkpoint expects + try: + raw_x = [available_values[name] for name in meta["input_features"]] + except KeyError as missing_feature: + raise KeyError( + f"GPR checkpoint expects input feature {missing_feature}," + " which is not available. Available features:" + f" {sorted(available_values.keys())}. Update the" + " 'available_values' mapping in this module to match the" + " checkpoint's input_features metadata." + ) from missing_feature + + raw_x = np.asarray([raw_x], dtype=float) + delta_mean, delta_std = predict_with_gpr_model(raw_x, model, likelihood) + # The noneccentric GPR predicts a correction (delta), to the PN approximation. + # This value is then added to the PN baseline to get the final corrected quantity. + corrected_value = baseline_value + float(delta_mean[0]) + logger.debug( + f"GPR correction for {quantity_name}: baseline={baseline_value:g}," + f" delta={float(delta_mean[0]):g} +/- {float(delta_std[0]):g}," + f" corrected={corrected_value:g}" + ) + return corrected_value + + +def _initial_orbital_parameters_gpr( + mass_ratio, + dimensionless_spin_a, + dimensionless_spin_b, + eccentricity, + separation, + orbital_angular_velocity, + num_orbits, + time_to_merger, + gpr_checkpoints, +) -> Tuple[float, float, float]: + """Zero-eccentricity initial orbital parameters, PN baseline, and + GPR correction. + """ + assert eccentricity == 0.0, ( + "Initial orbital parameters from GPR can currently only be computed for" + " zero eccentricity." + ) + # Start from the baseline PN guess, which the GPR models are trained to correct + pn_separation, pn_omega, pn_adot = _initial_orbital_parameters_pn( + mass_ratio=mass_ratio, + dimensionless_spin_a=dimensionless_spin_a, + dimensionless_spin_b=dimensionless_spin_b, + eccentricity=eccentricity, + separation=separation, + orbital_angular_velocity=orbital_angular_velocity, + num_orbits=num_orbits, + time_to_merger=time_to_merger, + ) + # Feature names match SimulationSupport.gpr (see the GPR tutorial notebook + # for a detailed explanation of how to train, save, and load the GPR model + # with real data). Each checkpoint selects the subset of these values it + # was trained on via its 'input_features' metadata. The aligned-spin example + # checkpoints in Examples use 'initial_separation', 'initial_mass_ratio', + # 'initial_dimensionless_spin1_z', and 'initial_dimensionless_spin2_z'. + available_values = { + "initial_separation": pn_separation, + "initial_mass_ratio": mass_ratio, + "initial_dimensionless_spin1_x": dimensionless_spin_a[0], + "initial_dimensionless_spin1_y": dimensionless_spin_a[1], + "initial_dimensionless_spin1_z": dimensionless_spin_a[2], + "initial_dimensionless_spin2_x": dimensionless_spin_b[0], + "initial_dimensionless_spin2_y": dimensionless_spin_b[1], + "initial_dimensionless_spin2_z": dimensionless_spin_b[2], + "pn_guess_omega": pn_omega, + "pn_guess_adot": pn_adot, + } + + corrected = {"Omega0": pn_omega, "Adot0": pn_adot} + for quantity_name, checkpoint_path in gpr_checkpoints.items(): + if quantity_name not in corrected: + raise ValueError( + f"Unknown quantity `{quantity_name}` in `gpr_checkpoints`." + f" Expected one of {','.join(sorted(corrected))}." + ) + corrected[quantity_name] = _apply_gpr_correction( + quantity_name, + corrected[quantity_name], + available_values, + checkpoint_path, + ) + orbital_angular_velocity = corrected["Omega0"] + radial_expansion_velocity = corrected["Adot0"] + logger.info( + "Selected approximately circular orbit using GPR corrected PN guess." + f" D0={pn_separation:g}, Omega0={orbital_angular_velocity:g}," + f" Adot0={radial_expansion_velocity:g}." + ) + return pn_separation, orbital_angular_velocity, radial_expansion_velocity + + +# CLI +# The function can be imported and called from Python directly, or it can be called with the CLI. +@click.command( + name="initial-orbital-parameters", +) +@click.option( + "--mass-ratio", + "-q", + type=float, + required=True, + help=r"Mass ratio, q = M_A / M_B \ge 1, of the two black holes.", +) +@click.option( + "--dimensionless-spin-a", + nargs=3, + type=float, + required=True, + help=( + "Dimensionless spin vector of the larger black hole, for example," + "written as '--dimensionless-spin-a 0.0 0.1 0.1'." + ), +) +@click.option( + "--dimensionless-spin-b", + nargs=3, + type=float, + required=True, + help=( + "Dimensionless spin vector of the smaller black hole, for example," + " written as '--dimensionless-spin-b 0.0 0.0 0.1'." + ), +) +@click.option( + "--eccentricity", + "-e", + type=float, + required=True, + help="Desired orbital eccentricity.", +) +@click.option( + "--mean-anomaly-fraction", + type=float, + help=( + "Mean anomaly divided by 2pi (between 0 and 1). Required if" + " eccentricity is nonzero." + ), +) +@click.option( + "--separation", + "-D", + type=float, + help="Coordinate separation, D_0, between the black holes.", +) +@click.option( + "--orbital-angular-velocity", + "-w", + type=float, + help="Orbital angular velocity, Omega_0.", +) +@click.option( + "--num-orbits", + type=float, + help="Desired number of orbits until merger.", +) +@click.option("--time-to-merger", type=float, help="Desired time until merger.") +@click.option( + "--method", + type=click.Choice(["PN", "GPR"]), + default="PN", + show_default=True, + help=( + "Compute from PN or from GPR, which applies a learned" + " correction to the PN baseline." + ), +) +@click.option( + "--gpr-omega-checkpoint", + type=click.Path(exists=True, dir_okay=False, readable=True), + help="Path to a trained GPR checkpoint providing Omega0 corrections.", +) +@click.option( + "--gpr-adot-checkpoint", + type=click.Path(exists=True, dir_okay=False, readable=True), + help="Path to a trained GPR checkpoint providing Adot0 corrections.", +) +@click.option( + "--output-json", + is_flag=True, + help="Print the result as a JSON file instead of text.", +) +def initial_orbital_parameters_command( + mass_ratio, + dimensionless_spin_a, + dimensionless_spin_b, + eccentricity, + mean_anomaly_fraction, + separation, + orbital_angular_velocity, + num_orbits, + time_to_merger, + method, + gpr_omega_checkpoint, + gpr_adot_checkpoint, + output_json, +): + """Estimate the initial orbital parameters for a BBH evolution. + + Estimates the initial coordinate separation, D_0, orbital angular velocity, Omega_0, and + radial expansion velocity, adot_0, from a Post-Newtonian approximation, optionally + corrected by a trained Gaussian Process Regression (GPR) model. + + Specify the target eccentricity and either '--separation', '--orbital-angular-velocity', + '--num-orbits', or '--time-to-merger'. + """ + _rich_traceback_guard = True + + target_params = { + "MassRatio": mass_ratio, + "DimensionlessSpinA": list(dimensionless_spin_a), + "DimensionlessSpinB": list(dimensionless_spin_b), + "Eccentricity": eccentricity, + } + if mean_anomaly_fraction is not None: + target_params["MeanAnomalyFraction"] = mean_anomaly_fraction + if num_orbits is not None: + target_params["NumOrbits"] = num_orbits + if time_to_merger is not None: + target_params["TimeToMerger"] = time_to_merger + + gpr_checkpoints = None + if method == "GPR": + gpr_checkpoints = {} + if gpr_omega_checkpoint: + gpr_checkpoints["Omega0"] = gpr_omega_checkpoint + if gpr_adot_checkpoint: + gpr_checkpoints["Adot0"] = gpr_adot_checkpoint + if not gpr_checkpoints: + raise click.UsageError( + "'--method GPR' requires either '--gpr-omega-checkpoint'," + " '--gpr-adot-checkpoint', or both." + ) + + D0, Omega0, Adot0 = initial_orbital_parameters( + target_params, + separation=separation, + orbital_angular_velocity=orbital_angular_velocity, + method=method, + gpr_checkpoints=gpr_checkpoints, + ) + + if output_json: + print( + json.dumps({"D0": D0, "Omega0": Omega0, "Adot0": Adot0}, indent=2) + ) + else: + rich.print(f"D0 = {D0}") + rich.print(f"Omega0 = {Omega0}") + rich.print(f"Adot0 = {Adot0}") + + return D0, Omega0, Adot0 + + +if __name__ == "__main__": + initial_orbital_parameters_command(help_option_names=["-h", "--help"]) diff --git a/tests/EccentricityControl/Test_InitialOrbitalParameters.py b/tests/EccentricityControl/Test_InitialOrbitalParameters.py index 4f0ca26..66ff9f3 100644 --- a/tests/EccentricityControl/Test_InitialOrbitalParameters.py +++ b/tests/EccentricityControl/Test_InitialOrbitalParameters.py @@ -1,10 +1,15 @@ # Distributed under the MIT License. # See LICENSE.txt for details. +import json + import numpy.testing as npt +import pytest +from click.testing import CliRunner from SimulationSupport.EccentricityControl.InitialOrbitalParameters import ( initial_orbital_parameters, + initial_orbital_parameters_command, ) @@ -41,13 +46,6 @@ def test_initial_orbital_parameters(): ), [15.6060791015625, 0.015, -4.541705362753467e-05], ) - npt.assert_allclose( - initial_orbital_parameters( - target_params, - orbital_angular_velocity=0.015, - ), - [15.6060791015625, 0.015, -4.541705362753467e-05], - ) npt.assert_allclose( initial_orbital_parameters( {**target_params, "NumOrbits": 20}, @@ -60,3 +58,198 @@ def test_initial_orbital_parameters(): ), [16.1357421875, 0.01430025219917298, -3.9831982447244026e-05], ) + + +def test_initial_orbital_parameters_pn_requires_zero_eccentricity(): + # PN method only supports zero eccentricity + target_params = { + "MassRatio": 1.0, + "MassA": 0.5, + "MassB": 0.5, + "DimensionlessSpinA": [0.0, 0.0, 0.0], + "DimensionlessSpinB": [0.0, 0.0, 0.0], + "Eccentricity": 0.1, + "MeanAnomalyFraction": 0.5, + } + with pytest.raises(AssertionError, match="zero eccentricity"): + initial_orbital_parameters( + target_params, + separation=16.0, + method="PN", + ) + + +def test_initial_orbital_parameters_gpr_requires_zero_eccentricity( + gpr_checkpoint_dir, +): + # GPR method currently only supports zero eccentricity + target_params = { + "MassRatio": 1.0, + "MassA": 0.5, + "MassB": 0.5, + "DimensionlessSpinA": [0.0, 0.0, 0.0], + "DimensionlessSpinB": [0.0, 0.0, 0.0], + "Eccentricity": 0.1, + "MeanAnomalyFraction": 0.5, + } + with pytest.raises(AssertionError, match="zero eccentricity"): + initial_orbital_parameters( + target_params, + separation=16.0, + method="GPR", + gpr_checkpoints={ + "Omega0": str(gpr_checkpoint_dir / "gpr_model_omega.pth") + }, + ) + + +def test_initial_orbital_parameters_gpr_requires_checkpoints(): + # method = "GPR" requires a non-empty gpr_checkpoints dict + target_params = { + "MassRatio": 1.0, + "MassA": 0.5, + "MassB": 0.5, + "DimensionlessSpinA": [0.0, 0.0, 0.0], + "DimensionlessSpinB": [0.0, 0.0, 0.0], + "Eccentricity": 0.0, + } + with pytest.raises(AssertionError, match="gpr_checkpoints"): + initial_orbital_parameters( + target_params, + separation=16.0, + method="GPR", + ) + + +def test_initial_orbital_parameters_gpr_rejects_unknown_quantities( + gpr_checkpoint_dir, +): + """ + Test that the keys of the checkpoint files match the parameter names used. + """ + target_params = { + "MassRatio": 1.0, + "MassA": 0.5, + "MassB": 0.5, + "DimensionlessSpinA": [0.0, 0.0, 0.0], + "DimensionlessSpinB": [0.0, 0.0, 0.0], + "Eccentricity": 0.0, + } + with pytest.raises(ValueError, match="Unknown quantity") as excinfo: + initial_orbital_parameters( + target_params, + separation=16.0, + method="GPR", + gpr_checkpoints={ + "omega": str(gpr_checkpoint_dir / "gpr_model_omega.pth") + }, + ) + assert "Omega0" in str(excinfo.value) + + +def test_initial_orbital_parameters_gpr_rejects_mismatched_checkpoint( + gpr_checkpoint_dir, +): + """ + Test that supplying a checkpoint trained for a different quantity is + prevented, rather than adding to the wrong correction. + """ + target_params = { + "MassRatio": 1.0, + "MassA": 0.5, + "MassB": 0.5, + "DimensionlessSpinA": [0.0, 0.0, 0.0], + "DimensionlessSpinB": [0.0, 0.0, 0.0], + "Eccentricity": 0.0, + } + with pytest.raises(ValueError, match="trained to predict"): + initial_orbital_parameters( + target_params, + separation=16.0, + method="GPR", + gpr_checkpoints={ + "Omega0": str(gpr_checkpoint_dir / "gpr_model_adot.pth") + }, + ) + + +def test_initial_orbital_parameters_gpr_omega_and_adot_correction( + gpr_checkpoint_dir, +): + """ + Test the GPR method with the real, trained checkpoints. The + expected deltas are computed directly from gpr_model_omega.pth + and gpr_model_adot.pth. + """ + target_params = { + "MassRatio": 1.0, + "MassA": 0.5, + "MassB": 0.5, + "DimensionlessSpinA": [0.0, 0.0, 0.0], + "DimensionlessSpinB": [0.0, 0.0, 0.0], + "Eccentricity": 0.0, + } + pn_separation, pn_omega, pn_adot = ( + 16.0, + 0.014474280975952748, + -4.117670632867514e-05, + ) + omega_delta = -3.0704395612701774e-05 + adot_delta = 8.766858081799e-05 + + separation, omega, adot = initial_orbital_parameters( + target_params, + separation=16.0, + method="GPR", + gpr_checkpoints={ + "Omega0": str(gpr_checkpoint_dir / "gpr_model_omega.pth"), + "Adot0": str(gpr_checkpoint_dir / "gpr_model_adot.pth"), + }, + ) + + npt.assert_allclose(separation, pn_separation) + npt.assert_allclose(omega, pn_omega + omega_delta, rtol=1e-4) + npt.assert_allclose(adot, pn_adot + adot_delta, rtol=1e-4) + + +def test_cli_gpr(gpr_checkpoint_dir): + runner = CliRunner() + result = runner.invoke( + initial_orbital_parameters_command, + [ + "--mass-ratio", + "1.0", + "--dimensionless-spin-a", + "0.0", + "0.0", + "0.0", + "--dimensionless-spin-b", + "0.0", + "0.0", + "0.0", + "--eccentricity", + "0.0", + "--separation", + "16.0", + "--method", + "GPR", + "--gpr-omega-checkpoint", + str(gpr_checkpoint_dir / "gpr_model_omega.pth"), + "--gpr-adot-checkpoint", + str(gpr_checkpoint_dir / "gpr_model_adot.pth"), + "--output-json", + ], + catch_exceptions=False, + ) + assert result.exit_code == 0, result.output + output = json.loads(result.output) + npt.assert_allclose( + output["Omega0"], + 0.014474280975952748 - 3.0704395612701774e-05, + rtol=1e-4, + ) + npt.assert_allclose( + output["Adot0"], + -4.117670632867514e-05 + 8.766858081799e-05, + rtol=1e-4, + ) diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..1d5fc1f --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,21 @@ +# Distributed under the MIT License. +# See LICENSE.txt for details. + +from pathlib import Path + +import pytest + +REPO_ROOT = Path(__file__).resolve().parent.parent + + +@pytest.fixture +def gpr_checkpoint_dir(): + """Directory containing example trained GPR checkpoints, + used for in Test_InitialOrbitalParameters.py.""" + return ( + REPO_ROOT + / "src" + / "SimulationSupport" + / "EccentricityControl" + / "Examples" + )